🏠 总目录📚 本教程 稳定性与容错
📑 本页目录(点开跳转)

21 · 训练稳定性与故障恢复

26 分钟 | ⭐ 千卡训练时,硬件故障是常态不是意外


🎯 一句话

在 1000 张 GPU 上跑一个月,"什么都不坏"的概率接近零。 大规模训练的工程重点不是"避免故障",而是"故障后能快速、正确地恢复" —— 而最可怕的不是崩溃,是那些不崩溃的错误。


📉 一、先接受一个数字

   假设单张 GPU 的平均无故障时间(MTBF)是 10000 小时

   1000 张卡同时跑:
   → 集群的 MTBF ≈ 10000 / 1000 = 10 小时 ⭐

   → 也就是说:【平均每 10 小时就会坏一次】
   → 一个月的训练会遇到几十次故障
   ⭐ 结论:checkpoint 和自动恢复不是"锦上添花",
     它决定了你的训练【能不能完成】

常见故障类型

故障 频率 特点
GPU 掉卡 / ECC 错误 会报错,好处理
网络抖动 / NCCL 超时 常常是挂死而不是报错 ⚠️
节点 OOM / 被抢占 云上竞价实例常见
磁盘满 / 慢 checkpoint 写不进去
⭐ 静默数据损坏 低但致命 不报错,训练照跑,结果是错的 💀

💾 二、Checkpoint 策略

   ⭐ 频率的权衡:

   太频繁 → 写 checkpoint 的时间占比高(70B 模型一次要几分钟)
   太稀疏 → 一次故障损失几小时的算力

   ⭐ 经验公式:
   最优间隔 ≈ √(2 × 单次checkpoint耗时 × MTBF)

   例:checkpoint 耗时 5 分钟,MTBF 10 小时
   → √(2 × 5 × 600) ≈ 77 分钟  → 每小时存一次差不多

四个必须做对的点

要点 说明
异步保存 先快速拷到 CPU 内存,后台再写盘 → 阻塞时间从几分钟降到几秒
分片保存 每个 rank 存自己那片(第 11 章),别汇总到 rank 0
原子写入 先写 .tmp 再 rename —— 防止写一半崩溃留下损坏文件
保留多份 至少留最近 3 个,防止最新的那个正好是坏的
# ⭐ PyTorch 分布式 checkpoint(异步 + 分片)
import torch.distributed.checkpoint as dcp

fut = dcp.async_save(
    {"model": model, "optim": opt, "step": step, "rng": get_rng_state()},
    checkpoint_id=f"ckpt/step-{step}")
# 训练继续,不阻塞 ⭐

⭐ checkpoint 里必须存什么

   □ 模型权重
   □ 优化器状态(m、v)
   □ 学习率调度器状态          ← 常被漏 ⚠️
   □ 当前 step / epoch
   □ 【数据加载器的位置】       ← 最常被漏 ⭐⭐
   □ 随机数状态(torch/numpy/python)
   □ 混合精度的 scaler 状态

   💥 漏了数据位置的后果:恢复后【重新从头读数据】
      → 某些样本被训练多次,某些完全没见过
      → 而这【不会报错】,你可能永远不知道

🔍 三、最可怕的:不崩溃的错误

   ⚠️ 崩溃是【好】故障——它会告诉你出事了
     真正危险的是【静默失败】:训练照跑,loss 也在降,但结果是错的

四种典型的静默失败

类型 表现 怎么发现
某张卡的梯度是坏的 loss 缓慢变差或停滞 定期校验各 rank 的梯度范数是否一致
数据管线错位 学到的东西不对 定期抽样打印实际喂进去的数据 ⭐
恢复后数据重复 过拟合特定样本 checkpoint 存数据位置
ECC 未纠正的位翻转 极偶发的数值异常 监控 ECC 计数器,nvidia-smi -q -d ECC
# ⭐ 一个便宜的静默错误探测器:定期检查各 rank 是否一致
import torch
def sanity_check_across_ranks(loss, grad_norm, step):
    t = torch.tensor([loss, grad_norm], device='cuda')
    gathered = [torch.zeros_like(t) for _ in range(dist.get_world_size())]
    dist.all_gather(gathered, t)
    if dist.get_rank() == 0:
        vals = torch.stack(gathered)
        spread = (vals.max(0).values - vals.min(0).values) / (vals.mean(0).abs() + 1e-9)
        if spread.max() > 0.5:          # ⭐ 各卡差异过大 = 有问题
            print(f"⚠️ step {step} rank 间差异异常: {vals}")

🔑 数据并行下各卡的 loss 应该不同(数据不同),但梯度范数在 AllReduce 后应该一致。 如果 AllReduce 之后各卡的梯度范数不一样 —— 通信出问题了。


🩹 四、训练不稳定:loss 尖峰

   大模型训练常见现象:loss 突然飙升,然后可能恢复、也可能再也回不来

   ⭐ 常见原因:
   ├─ 学习率过大 / warmup 太短
   ├─ 某个 batch 有异常数据(超长、乱码、重复)⭐
   ├─ FP16 溢出(换 BF16,第 6 章)
   ├─ 梯度裁剪阈值不合适
   └─ 数值不稳定(attention logits 过大、LayerNorm 位置)

五个实用的稳定化手段

手段 说明
梯度裁剪 clip_grad_norm_(params, 1.0) —— 几乎是标配 ⭐
跳过异常 batch grad_norm 超过历史中位数的 N 倍就跳过这一步
BF16 而非 FP16 避免溢出(第 6 章)⭐
充分的 warmup 大 batch 时尤其重要
回滚重跑 尖峰后不恢复 → 回到上一个 checkpoint,换个数据顺序重跑
# ⭐ 跳过异常 batch:非常便宜,收益很大
import torch
gn = torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
grad_history.append(gn.item())
median = statistics.median(grad_history[-100:])

if gn > 10 * median and len(grad_history) > 100:
    opt.zero_grad(set_to_none=True)        # ⭐ 丢弃这一步
    log(f"skip step {step}, grad_norm={gn:.1f} vs median {median:.1f}")
else:
    opt.step(); opt.zero_grad(set_to_none=True)

💡 "回滚 + 换数据顺序"是大模型训练的标准操作: 如果尖峰是某批数据引起的,换个顺序往往就跳过去了。 大厂的训练日志里这类回滚很常见 —— 它不是失败,是正常流程。


🔁 五、弹性训练与快速恢复

   ⭐ 目标:故障后【自动】恢复,不需要人半夜爬起来

   torchrun --nnodes=8 --nproc-per-node=8 \
            --max-restarts=10 \              # ⭐ 自动重启
            --rdzv-backend=c10d \
            --rdzv-endpoint=$MASTER:29500 \
            train.py

加快恢复的三个手段

手段 效果
checkpoint 放本地 NVMe 加载从几分钟降到几十秒 ⭐
热备节点 坏了立刻顶上,不用等调度
NCCL 超时设置合理 别让它挂死几十分钟才报错 ⭐
# ⭐ 三个必设的 NCCL 环境变量
export TORCH_NCCL_BLOCKING_WAIT=1          # 超时报错而不是永久挂起
export TORCH_NCCL_ASYNC_ERROR_HANDLING=1   # 异常时干净退出
export NCCL_TIMEOUT=1800                   # 30 分钟(默认可能太长)

💥 不设 BLOCKING_WAIT 的后果: 一张卡挂了,其他 999 张卡会永远等下去 —— GPU 利用率显示 100%(在空转),但一步都没往前走。 你可能几小时后才发现。第 10 章提过这个坑)


🔗 和站内其他章的关系

相关的地方 这里的位置
第 11 章 分片保存 checkpoint 的正确做法
第 10 章 AllReduce 挂死 本章给出完整解法
第 6 章 FP16 溢出 loss 尖峰的原因之一
ML 基础第 11 章 单机版的训练调试
《模型上线之后》17 版本回溯与可复现 可复现的三个层次 "恢复出来的是不是同一条轨迹" —— checkpoint 里该存什么,那边有更完整的清单 ⭐
《Kaggle竞赛方法论》04 超参数调优与工程实践 随机种子与可复现性 单机版的种子清单;多卡下还要管每个 rank 的种子和数据位置

✅ 检查点

  1. 1000 张卡的集群 MTBF 大概是多少?这意味着什么?
  2. checkpoint 间隔的经验公式是什么?
  3. checkpoint 的四个要点是什么?为什么要原子写入?
  4. checkpoint 里最常被漏掉的是什么?漏了会怎样?
  5. 为什么说"崩溃是好故障"?
  6. 怎么便宜地探测静默错误?各卡的 loss 和梯度范数应该一致吗?
  7. loss 尖峰的常见原因?五个稳定化手段?
  8. "跳过异常 batch"怎么实现?
  9. 不设 TORCH_NCCL_BLOCKING_WAIT 会怎样?
👀 答案
  1. 单卡 MTBF 10000 小时的话,1000 卡集群约 10 小时。意味着平均每 10 小时坏一次,一个月训练会遇到几十次故障——checkpoint 和自动恢复决定训练能否完成。
  2. √(2 × 单次checkpoint耗时 × MTBF)。例:耗时 5 分钟、MTBF 10 小时 → 约 77 分钟。
  3. 异步保存(先拷到 CPU 内存后台写盘)②分片保存(别汇总到 rank 0)③原子写入 ④保留多份。原子写入(先写 .tmp 再 rename)是为了防止写一半崩溃留下损坏文件
  4. 数据加载器的位置。漏了会导致恢复后重新从头读数据,某些样本训练多次、某些完全没见过,而且不会报错
  5. 因为崩溃会告诉你出事了。真正危险的是静默失败——训练照跑、loss 也在降,但结果是错的。
  6. 定期 all_gather 各 rank 的 loss 和梯度范数,检查差异loss 应该不同(数据不同),但梯度范数在 AllReduce 之后应该一致——不一致说明通信出问题了。
  7. 原因:学习率过大/warmup 太短、某个 batch 有异常数据、FP16 溢出、裁剪阈值不合适、数值不稳定。五手段:梯度裁剪跳过异常 batch用 BF16、充分 warmup、回滚 + 换数据顺序重跑
  8. 记录 grad_norm 历史,如果超过近 100 步中位数的 10 倍就 zero_grad() 丢弃这一步。非常便宜、收益很大。
  9. 一张卡挂了,其他卡会永远等下去——GPU 利用率显示 100%(空转)但一步都没前进,可能几小时后才发现。

🛑 可以停在这里

走神救援

千卡训练时硬件故障是常态:单卡 MTBF 10000 小时 → 1000 卡集群约 10 小时坏一次,一个月遇到几十次 → checkpoint 和自动恢复决定训练能不能完成。⭐间隔公式 √(2×checkpoint耗时×MTBF)(5 分钟耗时 + 10 小时 MTBF → 约每小时一次)。四个要点:⭐异步保存(先拷 CPU 再后台写盘,阻塞从几分钟降到几秒)、⭐分片保存(别汇总到 rank 0)、原子写入(先 .tmp 再 rename,防写一半崩溃)、留多份。⭐⭐checkpoint 里最常漏的是数据加载器位置——漏了会重新从头读数据,某些样本训多次某些没见过,而且不报错;还要存调度器状态、RNG、scaler。⭐最可怕的不是崩溃是静默失败(训练照跑 loss 也降但结果是错的)——崩溃是好故障,它会告诉你出事了;四种静默失败:某卡梯度坏了、数据管线错位、恢复后数据重复、ECC 位翻转。⭐便宜的探测器:定期 all_gather 各 rank 的 loss 和梯度范数——loss 应该不同(数据不同)但梯度范数在 AllReduce 后必须一致,不一致=通信出问题loss 尖峰的原因(学习率过大、异常数据、FP16 溢出、裁剪阈值、数值不稳定)和五个手段:梯度裁剪、⭐跳过异常 batch(grad_norm > 近100步中位数的 10 倍就丢弃这一步,极便宜收益大)、用 BF16、充分 warmup、⭐回滚+换数据顺序重跑这是大厂的标准操作不是失败)。弹性训练torchrun --max-restarts,checkpoint 放本地 NVMe、热备节点;⚠️⭐三个必设的 NCCL 环境变量——不设 TORCH_NCCL_BLOCKING_WAIT=1 的话一张卡挂了其他 999 张会永远等下去,GPU 显示 100% 但一步没走,几小时后才发现

下一节 👉 22-数据管线与存储.md

打卡记录保存在你的浏览器里,首页能看到总进度