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

PyTorch使用nn.MSELoss未指定target致损失停滞问题咨询

问题解答

原代码损失停滞的根本原因

首先明确nn.MSELoss的前向传播参数规则:它的forward方法签名为forward(input: Tensor, target: Tensor) -> Tensor,target没有设置任何默认值,不传的话会直接抛出参数缺失错误,不存在“未显式指定target就运行特殊逻辑”的情况。
你遇到的损失停滞问题和是否显式写target=关键字没有任何关系,核心错误是原代码对标签做了多余的unsqueeze(1)操作,导致预测值和标签形状不匹配,触发PyTorch广播机制后计算出的损失完全不符合预期。
最常见的场景是回归任务:假设batch size为N,模型输出y_pred形状为(N,),原始标签y_batch形状也为(N,),对y_batch做unsqueeze(1)后形状变为(N, 1)。广播机制会自动把两个张量扩展为(N, N)的形状再计算MSE——相当于每个样本的预测值,都和整个batch内所有N个样本的标签逐一计算了损失,此时反向传播的梯度方向完全混乱,训练到一定阶段损失就会卡在固定值不再下降。
你修改后的代码去掉了错误的unsqueeze(1)操作,传入的y_batch和y_pred形状完全匹配,损失计算逻辑回归正确,停滞问题自然解决。这里的target=关键字不影响运行结果,你可以试下去掉target=直接按位置传y_batch,训练效果不会有任何差异。

PyTorch损失函数的target传参规则

所有PyTorch内置损失函数都不强制要求显式以关键字形式指定target参数:

  • 所有内置损失的前向传播参数顺序是固定的:第一个位置参数为模型预测值input,第二个位置参数为真实标签target,按位置传参是完全合法的写法,官方文档的绝大多数示例也采用这种写法,比如MSELoss的官方示例就是直接按位置传参:loss = nn.MSELoss()(output, target)。
  • 只有当你不按默认顺序传参时,才需要用关键字明确指定参数对应关系。比如你非要把标签写在第一个参数位,就必须显式写input=y_pred标明预测值,否则会出现参数对应错误——MSE因为是对称损失可能不会立刻报错,但交叉熵这类非对称损失会直接导致训练崩溃。

调试提示:训练代码出现损失异常时,优先检查预测值和标签的形状是否匹配(分类任务中交叉熵等损失要求标签比预测值少一个类别维度的特殊场景除外),广播机制触发的隐式形状扩展是非常隐蔽的错误诱因,建议调试阶段随时打印张量形状确认。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.28 08:36:29