文章

LLM 训练 NaN/Inf 故障治理生产实战:用 Non-finite 门禁、GradScaler 证据与 Layer Bisect 隔离首个坏步

大模型训练中的 NaN/Inf 往往在多个步骤后才暴露。本文从非有限值门禁、GradScaler 跳步证据、分布式一致跳过、层级二分定位与坏批次隔离出发,给出可回放、可止损、可恢复的数值异常治理方案。

背景:NaN 往往不是故障起点,而是最后一张多米诺骨牌

大模型训练出现 NaN、+Inf 或 -Inf 时,日志里最先看到的通常是 loss=nan,但真正的故障可能发生在更早的位置:异常输入制造了极端 Logit,某层激活溢出,反向梯度在低精度下超出动态范围,随后非有限值进入梯度通信、优化器状态和模型参数。

因此,“检测到 NaN 后重启任务”不是治理方案。生产系统需要回答五个问题:

  1. 第一个非有限值出现在哪个 Step、Rank、Layer 和 Tensor?
  2. 坏值是否已经进入 Collective、Optimizer State 或检查点?
  3. 所有 Rank 是否对跳过更新作出了相同决定?
  4. 能否使用同一批次、同一随机状态和同一配置进行复现?
  5. 恢复后如何防止坏检查点重新进入训练链路?

这五个问题决定了数值异常是一次可诊断事件,还是一次只能从头重跑的事故。

核心原理:把数值异常拆成四个边界

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 的模型,可采用二分定位:

  1. 使用固定坏批次和固定 RNG 状态回放;
  2. 先检查中间层边界的激活与梯度是否有限;
  3. 判断异常位于前半区还是后半区;
  4. 逐轮缩小到具体 Block;
  5. 最后在目标 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 异常场景。

参考资料

  1. PyTorch, Automatic Mixed Precision examples
  2. PyTorch, Autograd anomaly detection
  3. PyTorch, clip_grad_norm_
  4. NVIDIA Megatron Core, DistributedDataParallelConfig
  5. NVIDIA Megatron FSDP, report_nan_in_param_grad
  6. PyTorch, torch.isfinite

常见问题

GradScaler 跳过一次 optimizer.step,是否说明训练已经损坏?
不一定。偶发溢出可能由动态 Loss Scale 探测引起,但必须记录发生频率、连续跳步次数和对应批次;连续或集中出现通常意味着输入、算子、学习率或精度配置存在问题。
发现 NaN 后,为什么不能只让当前 Rank 跳过?
数据并行中的所有 Rank 必须对本次更新做一致决定。某个 Rank 单独跳过,而其他 Rank 继续通信或更新,可能造成参数分叉、Collective 不匹配或后续恢复困难。
torch.autograd.detect_anomaly 能否长期在线开启?
不建议。它适合在已锁定的坏批次或短回放窗口中定位产生异常梯度的前向算子,长期启用会显著拖慢训练。
GradScaler 跳过一次参数更新后,Scheduler 是否应该前进?
通常应把 Scheduler 与成功的 Optimizer Step 对齐,而不是与每次训练循环对齐。若本轮参数没有更新却推进 Scheduler,会造成学习率计划和真实更新次数错位。