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

PyTorch中PowBackward0为何会引发异常NaN梯度?

问题描述

我有一个包含NaN的PyTorch张量,使用简单MSE Loss计算损失时,即使掩码去除NaN值,梯度仍会变为NaN。奇怪的是,仅当在计算含pow操作的损失后应用掩码时才会出现该问题。具体案例如下:

import torch
torch.autograd.set_detect_anomaly(True)

x = torch.rand(10, 10) 
y = torch.rand(10, 10)
w = torch.rand(10, 10, requires_grad=True)
y[y > 0.5] = torch.nan


o = w @ x
l = (y - o)**2
l = l[~y.isnan()]

try:
    l.mean().backward(retain_graph=True)
except RuntimeError:
    print('(y-o)**2 caused nan gradient')

l = (y - o)
l = l[~y.isnan()]

try:
    l.mean().backward(retain_graph=True)
except RuntimeError():
    pass
else:
    print('y-o does not cause nan gradient')

l = (y[~y.isnan()] - o[~y.isnan()])**2
l.mean().backward()
print('masking before pow does not propagate nan gradient')

请问为何经过pow函数的反向传播时,NaN梯度会发生传播?


问题解析

核心原因是NaN和0相乘的结果仍是NaN,结合PyTorch反向传播的计算逻辑,导致了梯度污染:

  • 先平方再掩码的情况

    1. 计算(y-o)**2时,y中的NaN会让对应位置的y-o变成NaN,平方后还是NaN。
    2. 反向传播时,平方操作的梯度是2*(y-o),这部分在y为NaN的位置会生成NaN值。
    3. 掩码操作的反向传播会给未选中的(原NaN)位置分配梯度0,此时就会触发NaN * 0的计算——结果还是NaN。
    4. 这些NaN值会被带入w的梯度计算流程,最终导致w的整体梯度变成NaN。
  • 先掩码再平方的情况
    先通过索引把NaN位置完全排除,再计算平方。整个过程没有NaN参与任何运算,反向传播时所有梯度都是正常数值,自然不会出现NaN梯度。

  • 直接计算y-o的情况
    虽然y-o存在NaN,但掩码操作后,反向传播时未选中位置的梯度被设为0,且没有平方操作带来的NaN*0计算。选中位置的梯度是正常的-1/N(N为有效样本数),因此w的梯度不会被污染。

简单说,先平方再掩码的操作,会在反向传播中产生NaN*0的无效计算;而先掩码再平方则从根源上避免了NaN参与运算,所以梯度正常。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.18 08:20:21