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

Pytorch Lightning训练AlexNet(CIFAR10)生成混淆矩阵异常求助

代码存在的核心错误如下:

  1. API混用错误
    你使用的是PyTorch生态的框架,但代码中所有张量操作、优化器调用都错误使用了TensorFlow的tf.前缀,这会直接导致运算逻辑完全异常,甚至输出张量的维度、数值全错,是混淆矩阵不符合预期的核心原因之一。你需要把所有tf.替换为torch.。
  2. AlexNet输出维度不匹配
    TorchVision提供的默认AlexNet模型最后一层输出维度为1000,适配ImageNet数据集的1000分类任务,而你做的是CIFAR10的10分类,没有修改最后一层的输出维度,会导致模型输出1000维的logits,损失计算逻辑完全错误,模型根本无法正常收敛。你需要在初始化时替换AlexNet的最后一层全连接层:
self.model = torchvision.models.alexnet(pretrained=False)
# 替换最后一层为10分类输出
self.model.classifier[6] = nn.Linear(self.model.classifier[6].in_features, num_classes)
  1. 混淆矩阵输入数据错误
    你存入OutputMatrix.predictions的是模型输出的原始10维logits张量,而生成混淆矩阵需要的是每个样本对应的预测类别ID,需要先对logits取维度1的最大值索引得到类别标签后再存储,修改如下:
# 取logits的最大值索引作为预测类别
OutputMatrix.predictions = torch.cat([tmp['preds'].argmax(dim=1) for tmp in outputs])
  1. 其他可优化点
  • PyTorch Lightning新版本不再推荐直接修改self.hparams,建议调用self.save_hyperparameters()统一管理超参数
  • 验证指标推荐使用self.log()方法进行记录,无需手动组装返回字典

内容的提问来源于stack exchange,提问作者nim_10

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.29 06:45:07