PyTorch混淆矩阵计算报样本数不一致错误排查
问题结论
你的判断完全正确。报错根源就是y_pred = y_pred.ge(.5).view(-1).cpu()这行的view(-1)操作,把原本形状为[482, 2]的预测张量(482个测试样本,每个样本对应A、B两个独立二分类任务的输出概率)展平成了长度964的一维张量,和长度为482的测试集标签样本数不匹配,最终触发样本数不一致的报错。
原因说明
- 原单二分类模型的输出形状为
[样本数, 1],此时用view(-1)展平后会得到长度等于样本数的一维预测数组,逻辑可以正常运行 - 改成双二分类输出后,每个样本对应2个独立的预测值,输出形状为
[样本数, 2],此时直接调用view(-1)会把两个任务的所有预测值打散拼接,总长度变为482*2=964,自然无法和482条样本的真实标签对齐。
修正方案
要得到形状为[482, 2]的0/1预测结果,只需要删除会破坏张量二维结构的view(-1)操作即可,修正代码如下:
y_pred = model(X_test) print(y_pred.shape) # torch.Size([482, 2]),对应482个样本的A、B两类预测概率 # 按0.5阈值转换为0/1预测值,保留[样本数, 2]的原始二维结构 y_pred = y_pred.ge(.5).cpu() print(y_pred.shape) # torch.Size([482, 2]),和标签形状完全匹配
后续计算注意事项
- 修正后得到的
y_pred是形状为[482, 2]的张量,第0列对应A类的二分类预测结果,第1列对应B类的二分类预测结果 - 如果需要分别计算两个任务的混淆矩阵,要单独取出对应列的预测值和对应任务的真实标签配对:比如计算A类混淆矩阵时取
y_pred[:, 0]和标签中A类的列配对,B类同理 - 如果使用sklearn的混淆矩阵接口,需要先把PyTorch张量转为numpy数组,在
.cpu()后追加.numpy()调用即可。
内容的提问来源于stack exchange,提问作者John David Deatherage
相关产品推荐
相关产品推荐

