PyTorch中如何通过布尔运算定位MNIST特定误分类图像索引?
解决PyTorch中特定误分类索引的布尔运算问题
报错原因
Python原生的and、or仅能处理单个布尔值,而test_predictions == 2和test_labels == 7返回的是布尔张量(每个元素对应一个样本的判断结果),直接用and会触发歧义错误——PyTorch无法将整个布尔张量判定为单一的True/False。
正确实现方式
针对张量的逻辑运算,需使用PyTorch提供的专用方法或运算符:
方法1:使用torch.logical_and()函数
这是PyTorch官方推荐的张量逻辑与运算方式:
# 生成"预测为2且真实标签为7"的布尔掩码 mask = torch.logical_and(test_predictions == 2, test_labels == 7) # 提取对应索引 wrong_indices = torch.where(mask)[0] # 若需去重(此处实际无需,因每个样本对应唯一索引) wrong_indices = wrong_indices.unique()
方法2:使用按位与运算符&
注意给每个比较表达式加括号(避免运算符优先级问题):
mask = (test_predictions == 2) & (test_labels == 7) wrong_indices = torch.where(mask)[0].unique()
扩展用法
针对混淆矩阵里的其他易混淆对(比如"预测为9但真实为4"),只需修改对应数值即可:
mask = torch.logical_and(test_predictions == 9, test_labels == 4) wrong_indices_4_9 = torch.where(mask)[0]
后续分析
拿到索引后,可提取对应图像和标签进行可视化分析:
# 假设test_images为测试集图像张量 wrong_images = test_images[wrong_indices] wrong_preds = test_predictions[wrong_indices] wrong_trues = test_labels[wrong_indices] # 可通过matplotlib等工具可视化,探究2和7手写风格相似等误分类原因
内容的提问来源于stack exchange,提问作者TennesseeTensor
相关产品推荐
相关产品推荐

