PyTorch训练ResNet时张量维度不匹配RuntimeError问题求助
问题分析与解决
核心错误原因
你的模型输出predictions形状是[batch_size, 4](每个样本对应4个旋转角度的预测得分),而真实标签rotation_label形状是[batch_size](每个样本对应一个角度的索引)。直接用==比较两个维度不匹配的张量,就会触发维度不匹配的RuntimeError。
从你给出的前5个预测结果能看到,每个样本有4个输出值;而标签是单个整数,两者维度完全不匹配。
修复方案
在计算正确预测数时,需要先对模型输出取argmax(dim=1),得到每个样本预测的角度索引(把[batch_size,4]的输出压缩成[batch_size]的类别索引),再和标签比较。
修改训练代码第28行:
# 原错误代码 curr_accum_correct += (predictions == rotation_label).sum().item() # 修改后代码 curr_accum_correct += (predictions.argmax(dim=1) == rotation_label).sum().item()
额外说明
- 如果你用的是
CrossEntropyLoss作为损失函数,训练时直接传入predictions和rotation_label是没问题的,因为PyTorch的CrossEntropyLoss会自动处理[N,C]的logits输入和[N]的标签输入。 - 只有在计算分类准确率时,才需要把模型的多类别输出转换成单类别索引,这一步是必须的。
内容的提问来源于stack exchange,提问作者sindhuja
相关产品推荐
相关产品推荐

