在深度学习模型训练过程中,数值不稳定是导致训练失败的主要原因之一。当我们计算损失函数、梯度或者归一化统计量时,如果某些中间结果超出了浮点数可表示的范围,或者参与了不合法的运算,就会得到NaN(Not a Number)。NaN具有传播性,一旦某个张量中出现NaN,后续所有依赖该张量的计算都会变成NaN,最终表现为loss变为NaN、准确率停滞或模型参数全部失效。理解NaN产生的底层机制,是快速定位并解决训练问题的第一步。
NaN产生的根本原因:溢出与除零
从计算机浮点数标准(IEEE 754)来看,单精度(float32)和半精度(float16)的数值范围都是有限的。以float32为例,其最大正规数约为3.4e38,超过该值就会发生上溢(overflow),结果变为inf。当inf参与减法(如inf减inf)或除以inf时,就会产生NaN。在神经网络中,过大的学习率、缺乏归一化的网络层、或者错误的损失函数实现,都可能让激活值或梯度在几步之内膨胀到溢出。
除零则是另一类直接引发NaN的操作。在批量归一化(BatchNorm)或层归一化(LayerNorm)中,我们计算方差后开根号,再用特征值减去均值再除以标准差。如果方差为零(例如某层输入全部相同),且代码中没有添加微小的平滑项,那么除以零就会得到NaN。类似的还有注意力机制中的softmax分母、对比学习中的相似度归一化等场景。
下面这段伪代码展示了一个典型的除零风险点:当输入张量所有元素相等时,std为零,直接除法导致NaN。
import torch
def unsafe_norm(x):
mean = x.mean(dim=-1, keepdim=True)
std = x.std(dim=-1, keepdim=True)
# 如果std为0,下面这行会产生NaN
return (x - mean) / std
x = torch.ones(4, 8) # 所有元素都是1,std=0
print(unsafe_norm(x)) # 包含NaN
使用钩子与断言定位首个异常张量
当训练在几百步之后才出现NaN,盲目调参效率极低。PyTorch提供的钩子(hook)机制允许我们在前向或反向传播时检查每一个模块的输出。通过注册前向钩子,我们可以打印或记录每一层输出的统计量(如最大值、是否含有NaN),从而精准找到第一个产生NaN的层。
除了钩子,在训练循环中加入断言(assert)也是一种轻量做法。例如在损失计算后立刻检查torch.isfinite(loss),若不为真则抛出错误并保存当前batch数据。配合梯度裁剪,我们可以在反向传播前拦截梯度中的inf。以下代码演示如何用钩子捕获NaN:
import torch
import torch.nn as nn
def nan_hook(module, inp, output):
if isinstance(output, torch.Tensor):
if not torch.isfinite(output).all():
print(f'NaN detected in {module.__class__.__name__}')
raise RuntimeError('NaN in output')
model = nn.Sequential(nn.Linear(10, 20), nn.ReLU(), nn.Linear(20, 2))
for name, mod in model.named_modules():
mod.register_forward_hook(nan_hook)
x = torch.randn(3, 10)
try:
model(x)
except RuntimeError as e:
print(e)
在实践中,建议将钩子仅用于调试阶段,因为频繁的张量检查会带来约5%到15%的性能开销。定位到问题层后,应针对性修改该层实现或输入处理逻辑,而非全局开启检查。
三类修复方案的原理与对比
针对溢出问题,梯度裁剪(gradient clipping)是最常用的手段。它通过限制梯度范数上限,防止参数更新步长过大引发激活爆炸。在PyTorch中可使用torch.nn.utils.clip_grad_norm_实现。该方法不改变前向数值分布,仅约束反向传播,因此适用性广,但无法解决前向除零问题。
混合精度训练(mixed precision)使用float16加速计算,但半精度动态范围极小,更容易溢出。此时需配合损失缩放(loss scaling):将损失乘以一个大常数,使梯度移入半精度可表示区间,反向后再缩放回来。若未正确设置缩放因子,反而会更频繁地出现NaN。相比之下,在归一化操作中添加epsilon(如1e-5)是消除除零最直接的方式,几乎所有官方实现都默认带此参数。
我们通过一个对比表格总结三种方案的特点:
| 方案 | 解决的主要问题 | 额外开销 | 使用建议 |
|---|---|---|---|
| 梯度裁剪 | 梯度溢出 | 低 | 默认开启,范数设1到10 |
| 损失缩放 | 半精度下溢与溢出 | 中 | 用AMP自动管理 |
| epsilon平滑 | 除零 | 无 | 所有除法分母必加 |
综合来看,防御NaN应当从前向数值安全(加epsilon、使用稳定激活函数如gelu替代不成熟自定义函数)和反向稳定(梯度裁剪、AMP)两端同时入手。每次修改训练脚本后,先用小数据跑五十步并开启NaN钩子验证,可大幅降低后期调试成本。
构建可复用的训练健康检查流程
为避免重复编写调试代码,可以将NaN检测封装为训练器的一部分。例如在每一个epoch开始前重置计数器,在backward之后扫描梯度,若发现非有限值则保存当前状态并退出。这种流程化设计让团队成员在接手项目时能立刻获得清晰的错误上下文,而不是面对一条毫无信息的NaN loss日志。
另一个常被忽略的点是数据预处理。输入特征中若含有缺失值(NaN)或极端离群点,会直接穿透到网络内部。因此在数据加载器中加入torch.nan_to_num或基于分位数的截断,能从源头消灭一类隐患。结合前文提到的钩子与裁剪策略,便可形成从数据、前向、反向到参数更新的完整防护链。
当模型结构复杂、包含自定义CUDA算子时,还需确认算子内部是否正确处理了边界情况。有些第三方算子未做分母保护,调用后即污染整个计算图。此时应在算子外包裹一层Python端的数值校验,或向维护者提交补丁。只有将检查意识融入工程习惯,才能彻底摆脱训练NaN的反复折磨。
NaNoverflowdivide_by_zero修改时间:2026-08-13 12:42:35