背景:NaN 往往不是故障起点,而是最后一张多米诺骨牌
大模型训练出现 NaN、+Inf 或 -Inf 时,日志里最先看到的通常是 loss=nan,但真正的故障可能发生在更早的位置:异常输入制造了极端 Logit,某层激活溢出,反向梯度在低精度下超出动态范围,随后非有限值进入梯度通信、优化器状态和模型参数。
因此,“检测到 NaN 后重启任务”不是治理方案。生产系统需要回答五个问题:
- 第一个非有限值出现在哪个 Step、Rank、Layer 和 Tensor?
- 坏值是否已经进入 Collective、Optimizer State 或检查点?
- 所有 Rank 是否对跳过更新作出了相同决定?
- 能否使用同一批次、同一随机状态和同一配置进行复现?
- 恢复后如何防止坏检查点重新进入训练链路?
这五个问题决定了数值异常是一次可诊断事件,还是一次只能从头重跑的事故。
核心原理:把数值异常拆成四个边界
1. 输入边界:先判断数据是否已经非有限
训练循环进入模型前,应对高风险字段进行轻量检查:
- 浮点输入是否包含 NaN/Inf;
- Token、Position ID、Mask 是否越界;
- Label 是否出现不允许的值;
- 样本长度、有效 Token 数和 Loss 权重是否异常;
- 自定义归一化、除法、对数、指数和概率计算的输入域是否合法。
不建议对每个大 Tensor 永久执行全量扫描。更实用的方法是:常态运行只检查输入摘要和关键标量;出现风险信号后,再对目标批次开启深度扫描。
2. 反向边界:在通信和参数更新前设置 Non-finite 门禁
loss 有限并不代表梯度有限。异常可能只在反向计算中产生,因此门禁应放在梯度通信或 Optimizer Step 之前。
PyTorch 的 clip_grad_norm_ 支持 error_if_nonfinite=True,可在总梯度范数为 NaN/Inf 时直接报错。Megatron Core 也提供在梯度 Collective 前检查 NaN/Inf 的配置。这类门禁的目标不是”把坏梯度裁剪成好梯度”,而是阻止坏值继续扩散。
3. 优化器边界:区分正常的 Loss Scale 探测和真实故障
在 FP16 混合精度训练中,GradScaler 会放大 Loss,反向后再还原梯度。如果发现非有限梯度,scaler.step(optimizer) 会跳过本次参数更新,并调整 Scale。
一次跳步可能只是动态 Scale 寻找合适区间,但以下情况应触发告警:
- 连续多个 Step 跳过;
- 跳步比例在短窗口内显著升高;
- Scale 持续下降但仍无法恢复;
- 同一数据分片或同一 Layer 反复触发;
- BF16/FP32 回放仍出现异常。
因此必须把 当前 Scale、前后 Scale、是否跳步、总梯度范数、学习率和批次指纹 一并写入证据包。
4. 分布式边界:跳过必须成为全局决策
数据并行训练中,一个 Rank 发现坏梯度后,不能只在本地 continue。建议每个 Rank 计算本地 bad_flag,然后使用一个小型 Collective 汇总;只要任意 Rank 报告异常,所有 Rank 都执行相同的动作:
- 不执行参数更新;
- 清空本轮梯度;
- 保存事件元数据;
- 按统一策略更新或冻结 Loss Scale;
- 进入回放、降级或终止分支。
这条规则尤其重要,因为非一致跳步会让不同 Rank 的参数和优化器状态发生分叉。
工程落地:构建”发现—止损—定位—恢复”闭环
第一步:记录首个坏步,而不是最后一次报错
为每个训练 Step 维护最小事件账本:
| 字段 | 说明 |
|---|---|
run_id / step / micro_step / global_batch_id | 训练定位 |
rank / dp_rank / tp_rank / pp_rank | 分布式定位 |
dataset_version / shard_id / sample_ids | 数据溯源 |
model_commit / tokenizer_hash / config_hash | 代码与配置指纹 |
precision / loss_scale / learning_rate | 精度与超参 |
total_grad_norm / skipped_step / first_bad_stage | 异常标志 |
rng_state_hash / checkpoint_parent | 复现与血缘 |
一旦发现非有限值,立即冻结当前事件记录。后续日志可以追加,但不能覆盖”首个坏步”。
第二步:按层级保存证据,不要直接 Dump 全模型
完整 Dump 所有激活和梯度会快速耗尽存储。更合理的是三级证据策略:
| 级别 | 内容 | 触发条件 |
|---|---|---|
| L1 常态证据 | Loss、Scale、Gradient Norm、输入摘要、配置指纹 | 每步 |
| L2 异常证据 | 首个异常 Rank、参数组、梯度桶、模块名称和 Tensor 统计 | 发现非有限值 |
| L3 定位证据 | 目标 Layer 的输入、输出、梯度和算子栈 | 坏批次回放时 |
Megatron FSDP 提供逐权重报告 NaN 梯度的调试能力,但官方文档也提示其性能代价较高。因此应把精细检测限制在复现窗口,而不是长期全量开启。
第三步:用 Layer Bisect 缩小异常范围
对包含大量 Transformer Block 的模型,可采用二分定位:
- 使用固定坏批次和固定 RNG 状态回放;
- 先检查中间层边界的激活与梯度是否有限;
- 判断异常位于前半区还是后半区;
- 逐轮缩小到具体 Block;
- 最后在目标 Block 内对 Attention、Normalization、MLP 和自定义算子加 Hook。
torch.autograd.detect_anomaly(check_nan=True) 能在调试时给出产生失败 Backward Function 的前向调用轨迹,但它会明显降低性能,所以适合坏批次的短窗口回放。
第四步:建立 Checkpoint 隔离规则
当某个 Step 出现异常时,不能默认它之前保存的所有检查点都安全。建议设置三类状态:
| 状态 | 含义 | 操作 |
|---|---|---|
| committed | 已通过前向、反向、梯度和参数有限性检查 | 允许恢复 |
| suspect | 位于首个坏步附近,尚未完成验证 | 需有限性扫描后放行 |
| quarantined | 明确包含非有限参数、优化器状态或无法复现的异常 | 禁止恢复 |
恢复任务只允许读取 committed 检查点。对 suspect 检查点,应执行参数、优化器状态和关键 Buffer 的有限性扫描后再决定是否放行。
PyTorch 参考实现:在更新前统一止损
下面的代码展示核心顺序。真实分布式系统还应加入 Rank 间 bad_flag 汇总和证据持久化。
import torch
def train_step(model, optimizer, scaler, batch, max_grad_norm: float):
optimizer.zero_grad(set_to_none=True)
with torch.autocast(device_type="cuda", dtype=torch.float16):
outputs = model(**batch)
loss = outputs.loss
if not torch.isfinite(loss):
raise FloatingPointError("non-finite loss before backward")
scaler.scale(loss).backward()
# 必须在读取、检查或裁剪真实梯度前进行 unscale
scaler.unscale_(optimizer)
try:
grad_norm = torch.nn.utils.clip_grad_norm_(
model.parameters(),
max_norm=max_grad_norm,
error_if_nonfinite=True,
)
except RuntimeError as exc:
optimizer.zero_grad(set_to_none=True)
# 在这里记录 batch 指纹、loss scale、rank 和配置指纹
raise FloatingPointError("non-finite gradients") from exc
scaler.step(optimizer)
scaler.update()
return {
"loss": float(loss.detach()),
"grad_norm": float(grad_norm.detach()),
"loss_scale": float(scaler.get_scale()),
}
关键顺序是:Backward → Unscale → 有限性检查/裁剪 → Step → Scale Update。不要在梯度仍处于缩放状态时使用其范数判断训练稳定性。
适用场景
这套方法适用于:
- FP16、BF16 或 FP8 混合精度预训练和微调;
- 使用 DDP、FSDP、Tensor Parallel 或 Pipeline Parallel 的分布式训练;
- 包含自定义 CUDA/Triton 算子、复杂损失函数或多模态输入的模型;
- 需要长时间运行、无法接受整轮重跑的大规模训练任务;
- 训练数据持续增量进入、可能混入极端样本的流水线。
常见误区
误区一:发现 NaN 后只降低学习率
学习率过高确实可能导致梯度爆炸,但 NaN 也可能来自非法输入域、错误 Mask、归一化分母为零、低精度溢出、损坏样本或自定义 Kernel。直接降学习率可能暂时掩盖问题,却没有消除根因。
误区二:梯度裁剪可以修复 NaN
裁剪只对有限梯度按比例缩放。梯度已经是 NaN/Inf 时,应拒绝更新并定位来源,而不是继续裁剪。
误区三:GradScaler 跳步后不需要记录
偶发跳步可以是正常现象,但没有证据就无法区分正常 Scale 探测和系统性数值故障。至少要记录 Scale、Step、Rank、批次和连续跳步次数。
误区四:只在 Rank 0 检查
异常可能只出现在某个数据并行 Rank 或流水线 Stage。只检查 Rank 0 会漏掉局部坏值,并可能让通信阶段把异常传播到其他设备。
误区五:开启 Anomaly Detection 后继续跑完整训练
Anomaly Detection 是诊断工具,不是常态监控方案。更合理的是先通过轻量门禁锁定坏批次,再用它做短窗口复现。
BF16 是否可以取消 Non-finite 门禁?
不可以。BF16 相比 FP16 更不容易因动态范围不足而溢出,但非法数学操作、异常输入、梯度爆炸和自定义 Kernel 错误仍会产生 NaN/Inf。
是否需要扫描全部参数以确认 Checkpoint 安全?
常态下可以使用分片级摘要、梯度桶检查和保存前门禁;出现异常后,对临近检查点执行完整有限性扫描更稳妥。逐参数扫描应控制频率,避免显著影响训练吞吐。
上线检查清单
- 输入、Loss、梯度和参数更新前都有明确的有限性检查边界。
- GradScaler 的 Scale、跳步结果和连续跳步次数已纳入指标。
- 所有 Rank 对异常 Step 的跳过、终止或恢复动作保持一致。
- 首个坏步能关联数据版本、样本 ID、配置、代码和父检查点。
- 支持使用固定 Batch、RNG State 和环境指纹进行回放。
- Layer Bisect 和目标模块 Hook 有开关,不会常态拖慢训练。
- 可疑和损坏 Checkpoint 不会被自动恢复流程选中。
- 已演练输入 NaN、梯度 Inf、连续跳步和单 Rank 异常场景。
参考资料
- PyTorch, Automatic Mixed Precision examples
- PyTorch, Autograd anomaly detection
- PyTorch, clip_grad_norm_
- NVIDIA Megatron Core, DistributedDataParallelConfig
- NVIDIA Megatron FSDP, report_nan_in_param_grad
- PyTorch, torch.isfinite