绝大多数情况下loss.backward()后梯度变nan是因中间张量触发非法运算(如log(0)、sqrt(负数)),而非模型结构错误;使用torch.autograd.set_detect_anomaly(true)可精准定位异常操作位置。

为什么loss.backward()后梯度变成NaN
绝大多数情况不是模型写错了,而是某个中间张量在前向或反向过程中触发了非法运算:比如log(0)、0 * inf、inf / inf,或者torch.sqrt()输入负数。这些操作不会立即报错,但会把结果设为nan,并污染后续所有梯度。尤其当使用nn.CrossEntropyLoss时,内部log_softmax对全-inf的logits输入会直接产出nan;TransformerEncoder中若src_key_padding_mask构造不当,也可能让softmax输入全为-inf,导致输出全nan。
用torch.autograd.set_detect_anomaly(True)定位源头
这是最直接有效的调试手段——它会让backward()在遇到nan梯度时立刻抛出带完整调用栈的异常,精确到哪一层、哪个操作。
使用方式很简单:
torch.autograd.set_detect_anomaly(True) loss.backward() # 这里会中断并打印出错位置
注意两点:
SkillSub Pro - Python 题解与代码注释双功能技能功能概述SkillSub Pro - Python 题解与代码注释双功能技能是一项面向实际任务的技能,主要用于SkillSub Pro 是一个 Python 题解生成与代码注释的 双功能合体技能 ,专为学生、算法学习者和开发者设计;✅ 一个技能,两种用途 :;核心要点📝 题解模式 :输入题目/题号,自动生成完整 Python 题解(含详细注释、解题思路、复杂度分析);💬 注释模式 :输入 Python 代码,自动添加详细中。它将相关步骤、
- 必须在
loss.backward()之前设置,且只设一次即可 - 开启后训练速度明显下降,仅用于调试,不要留在生产训练循环中
- 如果报错指向某层
forward,说明问题出在该层输入(比如Linear前的logits已含nan)
检查梯度范数和中间变量是否含nan
在训练循环中加几行轻量检查,比等整个训练崩掉再回头查快得多:
- 每次
backward()后立刻检查:if torch.isnan(loss).any(): print("loss nan at step", step); break - 在关键节点插入断言:
assert not torch.isnan(x).any(), f"x is nan before {layer_name}",例如softmax前、loss输入前、Linear输出后 - 打印梯度范数:
print_gradient_norms(model)函数能帮你快速识别哪层梯度最先失控(比如某RNN层梯度从0.5跳到1e5)
哪些操作最容易悄悄引入nan?
以下场景高频触发,需逐项排查:
-
torch.log()或torch.log1p()输入≤0:务必先x.clamp(min=1e-8),尤其在自定义loss或概率归一化后 - 学习率过大:尝试降10倍(如
1e-3 → 1e-4),哪怕用Adam也别忽略lr超参 -
BatchNorm分母为0:小batch或全同特征时方差为0,eps默认1e-5不够用,可试1e-3 - 混合精度训练未配
GradScaler:必须用scaler.step(optimizer)而非optimizer.step(),且scaler.update()不能漏 -
Transformer中src_key_padding_mask填错:掩码值应为True(mask)或False(keep),误填0/1或nan会破坏softmax数值稳定性
真正麻烦的是nan被mean()或sum()掩盖后延迟爆发——可能第100步才暴露,但污染从第2步就开始了。
Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!










