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反向传播的计算逻辑,导致了梯度污染:
先平方再掩码的情况
- 计算
(y-o)**2时,y中的NaN会让对应位置的y-o变成NaN,平方后还是NaN。 - 反向传播时,平方操作的梯度是
2*(y-o),这部分在y为NaN的位置会生成NaN值。 - 掩码操作的反向传播会给未选中的(原NaN)位置分配梯度0,此时就会触发
NaN * 0的计算——结果还是NaN。 - 这些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
相关产品推荐
相关产品推荐

