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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.30 19:48:02