MNIST项目中测试准确率与混淆矩阵结果不匹配问题排查
MNIST模型准确率与混淆矩阵不匹配的问题排查与修复
核心问题定位
1. 模型最后一层错误使用ReLU激活
分类任务中,模型输出的logits无需额外激活(如CrossEntropyLoss会自动处理logits),但你的模型最后一层添加了ReLU,会将负数输出置为0,完全破坏类别间的概率分布,导致预测结果失真,这也是非CNN模型却出现过高准确率的异常原因之一。
修复代码:
修改模型的层堆叠,移除最后一层的ReLU:
self.layer_stack = nn.Sequential( nn.Flatten(), nn.Linear(in_features=input_shape, out_features=hidden_units), nn.ReLU(), nn.Linear(in_features=hidden_units, out_features=output_shape) # 移除最后的ReLU )
2. Softmax维度参数设置错误
生成预测标签时,torch.softmax(pred_logits.squeeze(), dim=0)的维度选择完全错误:
- pred_logits的形状为
[batch_size, num_classes],Softmax应在**类别维度(dim=1)**计算,保证每个样本的类别概率和为1; - 当前用dim=0是对整个batch的所有类别做Softmax,导致概率分布完全混乱,argmax得到的预测标签毫无意义,直接造成混淆矩阵计算的准确率偏低。
修复代码:
调整Softmax的维度参数为1,同时可移除多余的squeeze()(若pred_logits无冗余维度):
y_pred = torch.softmax(pred_logits, dim=1).argmax(dim=1)
3. 评估函数中数据未同步设备(潜在问题)
你的eval_model函数未将X和y移动到模型所在设备:
- 若模型在GPU、数据在CPU,虽PyTorch会自动将X移至GPU,但y仍在CPU,计算准确率时会隐式同步,可能引发效率问题或潜在错误。
修复代码:
在eval_model的循环中添加设备同步:
for X,y in tqdm(data_loader): X, y = X.to(device), y.to(device) # 新增设备同步代码 y_pred = model(X) loss += loss_fn(y_pred, y) acc += accuracy_fn(y_true=y, y_pred=y_pred.argmax(dim=1))
验证修复效果
修复上述问题后重新训练模型:
- 模型测试准确率会回归合理范围(非CNN模型在MNIST上达到97%-98%属于正常水平);
- 混淆矩阵对角线求和得出的准确率会与评估函数输出的准确率完全匹配。
内容的提问来源于stack exchange,提问作者Nico
相关产品推荐
相关产品推荐

