浮点数精度误差本身不会直接让模型崩溃,但当它出现在梯度计算、损失函数分母、归一化条件判断或early stopping阈值中时,会触发nan或inf的连锁反应。

为什么 0.1 + 0.2 != 0.3 会让模型训练突然崩溃?
浮点数精度误差本身不会直接让模型“崩溃”,但当它出现在梯度计算、损失函数分母、归一化条件判断或 early stopping 阈值里时,就会触发 NaN 或 inf 的连锁反应。比如 torch.sqrt(-1e-18) 算出 nan,再反向传播就全毁了。
- 常见诱因:手动写
1e-8做 epsilon 保护,但没考虑实际数值范围(如在1e-20量级数据上仍用1e-8) - 容易被忽略的场景:用
np.array([0.1, 0.2, 0.3]).sum() == 0.6做逻辑分支——结果是False,后续代码跳过关键初始化 - PyTorch/TensorFlow 默认 float32 下,相对误差约
1e-7;float64 能压到1e-16,但显存翻倍、速度下降
用 torch.finfo 和 np.finfo 动态查 epsilon,别硬编码 1e-8
固定 epsilon 在不同数据尺度下会失效:对大数值(如图像像素值 0–255)1e-8 太小,对小梯度(如深层网络末层 1e-12)又太大。
- PyTorch 推荐写法:
eps = torch.finfo(x.dtype).tiny(返回最小正正规数)或.eps(机器精度) - NumPy 对应:
np.finfo(x.dtype).smallest_subnormal或.resolution - 实际例子:在 LayerNorm 分母中写成
x / (var.sqrt() + torch.finfo(var.dtype).tiny),比+ 1e-5更鲁棒
避免用 == 比较浮点数,改用 torch.allclose 或 np.isclose
直接比较两个 tensor 是否相等,哪怕只差一个 ulp(unit in last place),也会返回 False,导致训练循环意外退出或 checkpoint 被跳过。
- 安全写法:
torch.allclose(a, b, atol=1e-8, rtol=1e-5)—— 同时检查绝对误差和相对误差 - 注意
rtol默认是1e-5,如果 a/b 接近 0,必须显式设atol,否则allclose可能误判 - 不要用
np.array_equal或torch.equal做数值相等判断,它们不做容差处理
训练中检测 NaN 和 inf 的低成本方法
等 loss 变成 nan 再报错,往往已经错过最早出问题的 layer。得在前向/反向关键节点主动拦截。
- PyTorch:在 loss.backward() 前加
assert torch.isfinite(loss).all(), f"Loss is {loss}" - 更细粒度:在 model output 后插一句
assert torch.isfinite(out).all(),快速定位爆炸源头 - 避免用
torch.isnan(x).any()—— 它比torch.isfinite(x).all()慢约 30%,且不捕捉inf
isfinite 替代 isnan、把 epsilon 当变量而非常量。最危险的不是误差本身,是它藏在条件分支里无声地拐弯。Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!











