Pytorch Lightning训练AlexNet(CIFAR10)生成混淆矩阵异常求助
代码存在的核心错误如下:
- API混用错误
你使用的是PyTorch生态的框架,但代码中所有张量操作、优化器调用都错误使用了TensorFlow的tf.前缀,这会直接导致运算逻辑完全异常,甚至输出张量的维度、数值全错,是混淆矩阵不符合预期的核心原因之一。你需要把所有tf.替换为torch.。 - 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)
- 混淆矩阵输入数据错误
你存入OutputMatrix.predictions的是模型输出的原始10维logits张量,而生成混淆矩阵需要的是每个样本对应的预测类别ID,需要先对logits取维度1的最大值索引得到类别标签后再存储,修改如下:
# 取logits的最大值索引作为预测类别 OutputMatrix.predictions = torch.cat([tmp['preds'].argmax(dim=1) for tmp in outputs])
- 其他可优化点
- PyTorch Lightning新版本不再推荐直接修改
self.hparams,建议调用self.save_hyperparameters()统一管理超参数 - 验证指标推荐使用
self.log()方法进行记录,无需手动组装返回字典
内容的提问来源于stack exchange,提问作者nim_10
相关产品推荐
相关产品推荐

