You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.07 00:22:07