nan损失主因是数值溢出或非法运算,如log(0)、inf/inf等;定位可用torch.autograd.set_detect_anomaly(true);高频场景包括log输入≤0、学习率过大、batchnorm方差为0、amp未正确缩放。

为什么训练中突然出现 nan 损失?
绝大多数情况不是模型“坏了”,而是数值溢出或非法运算在某个环节悄悄发生。最常见的是 log(0)、inf / inf、0 * inf 这类操作,它们会直接污染梯度,让后续所有 loss 变成 nan。尤其在使用 nn.CrossEntropyLoss 或 nn.BCEWithLogitsLoss 时,如果输入 logits 出现极端值(比如全 -inf 或极大正数),内部 log_softmax 或 sigmoid 就可能触发 nan。
如何快速定位 nan 源头?
别靠猜,用 PyTorch 内置的梯度检查工具直接打断点:
- 在
loss.backward()前加torch.autograd.set_detect_anomaly(True),它会在反向传播遇到nan时抛出带栈追踪的异常,明确指出哪一层输出/梯度出问题 - 训练循环中定期检查:
if torch.isnan(loss): print("loss is nan at step", step); break - 更细粒度地检查中间变量:
assert not torch.isnan(x).any(), f"x contains nan at {layer_name}",插在关键层(如 softmax 前、loss 输入前)
哪些操作最容易引入 nan?
以下场景高频踩坑,需逐项排查:
-
torch.log()或torch.log1p()输入 ≤ 0:确保输入先做clamp(min=1e-8),尤其在自定义 loss 或概率归一化后 - 学习率过大导致权重爆炸:尝试把
lr降低 10 倍(例如从1e-3改为1e-4),观察是否消失 - BatchNorm 层输入方差为 0(单样本 batch 或全零输入):检查 dataloader 是否混入异常样本,或临时改用
nn.LayerNorm - 混合精度训练(
amp)中未正确处理inf梯度:务必启用scaler.step(optimizer)而非直接optimizer.step(),且scaler.update()不可省略
训练中实时防御 nan 的实用技巧
预防比调试省力,几个低成本加固手段:
- 在 optimizer.step() 前加梯度裁剪:
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0),避免梯度爆炸引发后续nan - loss 计算后立刻检查:
if torch.isfinite(loss).all(): optimizer.step() else: print("skip step due to nan loss"); optimizer.zero_grad() - 初始化权重时避开极端值:用
torch.nn.init.xavier_normal_(m.weight)替代默认初始化,尤其对深层网络有效 - 验证集评估时也检查 loss —— 有时
nan只在 eval 模式下暴露(比如 dropout 关闭后某路径激活)
真正麻烦的不是第一次出现 nan,而是它被 mean() 或 sum() 掩盖后延迟爆发。所以检查要趁早,别等 loss 曲线断崖才动手。
Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!











