跳转至

训练稳定性:Loss Spike、NaN、梯度异常诊断

更新日期:2026-04-14


一、训练失败模式

失败模式 现象 根因 典型修复
Loss Spike loss 突涨 2-10× 数据异常 / 梯度爆炸 / lr 过大 跳过 batch / clip / lr 降
NaN loss / grad → NaN FP16 上溢、除零、log(0) FP32 master / loss scale
梯度消失 grad norm → 0 Pre-LN 配错、残差路径断 改架构 / 加 skip
梯度爆炸 grad norm → ∞ 数据病态 / 初始化错 clip + 重新 warmup
收敛停滞 loss 长时间不降 数据重复 / 优化器停滞 数据洗牌 + 检查 optimizer state
Catastrophic forgetting 后期评测下降 lr 太大 / annealing 数据有偏 cooldown 用高质量数据
Loss divergence loss 单调上升 TP 配错 / 通信 bug 检查并行配置 + nccl

二、Loss Spike 诊断与修复

2.1 什么是 Loss Spike

训练过程中 loss 突然跳高(通常 2-10x),然后可能自行恢复或不恢复。

2.2 常见原因

2.3 实时监控脚本

class TrainingMonitor:
    def __init__(self, window=100, spike_threshold=2.0):
        self.loss_history = []
        self.grad_norm_history = []
        self.window = window
        self.spike_threshold = spike_threshold

    def check(self, step, loss, grad_norm, lr, batch_info=None):
        self.loss_history.append(loss)
        self.grad_norm_history.append(grad_norm)

        if len(self.loss_history) < self.window:
            return

        recent_avg = mean(self.loss_history[-self.window:])

        # Loss spike 检测
        if loss > recent_avg * self.spike_threshold:
            print(f"LOSS SPIKE at step {step}!")
            print(f"  loss={loss:.4f}, recent_avg={recent_avg:.4f}")
            print(f"  grad_norm={grad_norm:.4f}")
            print(f"  lr={lr:.2e}")
            if batch_info:
                print(f"  batch_info={batch_info}")
            # 保存 checkpoint 用于事后分析
            save_debug_checkpoint(step)

        # Grad norm 异常检测
        recent_grad_avg = mean(self.grad_norm_history[-self.window:])
        if grad_norm > recent_grad_avg * 5:
            print(f"GRAD NORM SPIKE at step {step}!")
            print(f"  grad_norm={grad_norm:.4f}, avg={recent_grad_avg:.4f}")

        # NaN 检测
        if isnan(loss) or isinf(loss):
            print(f"NaN/Inf DETECTED at step {step}!")
            save_debug_checkpoint(step)
            raise RuntimeError("Training diverged")

三、NaN 诊断

3.1 NaN 的来源

3.2 NaN 定位方法

# 方法 1: 逐层检测
def check_nan_hook(module, input, output):
    if isinstance(output, torch.Tensor):
        if torch.isnan(output).any():
            print(f"NaN in {module.__class__.__name__}")
            print(f"  input nan: {torch.isnan(input[0]).any()}")
            raise RuntimeError(f"NaN detected in {module.__class__.__name__}")

for name, module in model.named_modules():
    module.register_forward_hook(check_nan_hook)

# 方法 2: torch.autograd.detect_anomaly()
with torch.autograd.detect_anomaly():
    loss = model(batch)
    loss.backward()
# 会在反向传播遇到 NaN 时打印完整 traceback
# 但速度极慢! 只用于 debug

四、初始化与训练早期稳定性

4.1 权重初始化

输出投影的特殊初始化很关键:每个残差块的输出投影 std 需要缩小 1/sqrt(2*n_layers)。这确保残差连接的方差不随深度爆炸。

def init_weights(model, n_layers):
    for name, param in model.named_parameters():
        if 'embed' in name:
            nn.init.normal_(param, std=0.02)
        elif 'W_o' in name or 'W_down' in name:
            # 输出投影: 缩小初始化
            nn.init.normal_(param, std=0.02 / sqrt(2 * n_layers))
        elif param.ndim >= 2:
            nn.init.normal_(param, std=0.02)
        else:
            nn.init.zeros_(param)  # bias

4.2 Warmup 的重要性

不使用 warmup 直接用高 lr 训练几乎一定会导致前几百步 loss spike 或 NaN。 Warmup 的作用: 训练初期,Adam 的二阶动量 v 还没有足够的统计量。当 v ≈ 0 时,自适应学习率 m/sqrt(v) 会变得极大,导致参数更新过大。Warmup 让学习率从 0 缓慢增加,给 v 时间累积足够的统计信息。


五、万卡训练的特殊问题

5.1 慢节点 (Straggler)

5.2 弹性训练与容错

万卡训练中,节点故障是必然事件。以 2000 GPU 集群为例,单 GPU MTBF 约 1000 小时,则集群级 MTBF 约 0.5 小时——即平均每 30 分钟就有一个 GPU 出问题。

解决方案 1:快速 checkpoint + 重启

  • 异步 checkpoint:不阻塞训练

  • 增量 checkpoint:只保存变化的部分

  • 间隔:每 10-30 分钟保存一次

解决方案 2:弹性训练

  • PyTorch Elastic (torchrun):节点故障后自动重新分配

  • 但 Megatron 的 TP/PP 需要所有 rank 存在,弹性有限

解决方案 3:冗余计算

  • 关键 rank 做冗余计算,故障时切换到备份

  • 成本增加 5-10%,但避免了任何停机

5.3 Checkpoint 策略


六、MoE 特有的稳定性问题

问题 表现 修复
路由崩塌 少数专家吃大部分 token Aux-Loss-Free + z-loss
专家梯度方差大 不同专家的梯度 norm 差异 100x per-expert gradient clipping
All-to-All 超时 负载不均导致等待 capacity factor 限制
训练后期路由退化 早期好的路由分布后期崩塌 监控路由熵, 发现下降时干预

参考文献

  • [1] Chowdhery et al. PaLM: Scaling Language Modeling with Pathways (包含训练稳定性讨论). 2022. 论文

  • [2] ML Engineering: Training Performance 指南

  • [3] Zhang et al. OPT: Open Pre-trained Transformer Language Models (训练日志公开). 2022. 论文

  • [4] Glorot & Bengio. Understanding Difficulty of Training Deep Feedforward Networks. 2010. 论文


上级 · C. 分布式训练基础设施