PyTorch使用MSE损失时出现张量形状不匹配RuntimeError
问题分析与解决
核心问题
看错误栈里的代码行——你在第118行已经正确用labels_onehot计算损失,但第121行又重复调用了criterion(outputs, labels),这里的labels是原始类别索引张量(形状应为[64],对应batch size=64),而模型输出outputs的形状是[64,10],两者维度不匹配,导致MSE损失计算时广播失败,抛出维度不匹配的错误。
你打印的是labels_onehot的形状,它和outputs形状一致,但实际报错的是用原始labels的那行代码,这就是问题根源。
解决步骤
- 删除/注释错误代码行:直接去掉第121行的
loss = criterion(outputs, labels),保留第118行的loss = criterion(outputs, labels_onehot)即可。 - 确认形状匹配:确保
outputs的形状为[batch_size, num_classes],labels_onehot通过nn.functional.one_hot(labels, num_classes=10).float()生成后,形状也应为[batch_size, 10],两者维度完全对应才能正常计算MSE损失。
额外注意
用MSE损失替代交叉熵且移除softmax层时,模型输出不需要做归一化处理,直接和one-hot格式的标签计算均方误差即可,这是一种回归式的分类思路,和交叉熵依赖概率分布的逻辑不同。
内容的提问来源于stack exchange,提问作者noobie
相关产品推荐
相关产品推荐

