PyTorch迁移学习绘制混淆矩阵类型报错问题求助
错误原因
- 张量初始化类型错误:你用来存储预测结果和真实标签的
all_preds、source_value张量初始化时没有指定类型,PyTorch默认生成float32类型的张量,后续拼接的整数值会被自动转为float,用float值做数组索引就会触发IndexError。 - 张量操作逻辑错误:你存储的
all_preds已经是max(1, keepdim=True)得到的类别索引,不需要再调用argmax(dim=1),多此一举会导致预测结果全部错误。同时你尝试用int()直接转换多元素张量为整数,只有单元素张量支持转Python标量,所以触发ValueError。
修复代码
1. 修改存储张量的初始化
将train_alexnet函数中原来的初始化代码:
all_preds = torch.tensor([]) source_value = torch.tensor([])
修改为:
# 显式指定为长整型,和标签类型一致,同时匹配计算设备 all_preds = torch.tensor([], dtype=torch.long, device='cuda' if cuda else 'cpu') source_value = torch.tensor([], dtype=torch.long, device='cuda' if cuda else 'cpu')
2. 修正堆叠张量的逻辑
删除多余的argmax操作,调整为:
# all_preds是[样本数, 1]维度,squeeze转为一维和真实标签匹配 stacked = torch.stack( (source_value, all_preds.squeeze()), dim=1 )
3. 简化混淆矩阵生成逻辑
直接用sklearn提供的方法生成混淆矩阵,避免手动循环的类型问题:
with torch.no_grad(): # 把张量从GPU移到CPU转numpy数组 y_true = source_value.cpu().numpy() y_pred = all_preds.squeeze().cpu().numpy() cmt = confusion_matrix(y_true, y_pred) print("混淆矩阵:\n", cmt)
测试集混淆矩阵修改说明
test_alexnet函数中生成混淆矩阵的逻辑和上面一致,只需要对应修改存储标签的张量初始化、删除多余的argmax操作即可。
内容的提问来源于stack exchange,提问作者hot_rod
相关产品推荐
相关产品推荐

