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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.26 10:20:04