训练稳定性: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. 分布式训练基础设施