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

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 个,防止最新的那个正好是坏的
# 🧩 骨架:`model` 来自你自己的代码,这一段只看写法
# ⭐ 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 尖峰

关键信息

五个实用的稳定化手段:

手段 说明
梯度裁剪 clip_grad_norm_(params, 1.0) —— 几乎是标配 ⭐
跳过异常 batch ⭐ grad_norm 超过历史中位数的 N 倍就跳过这一步
BF16 而非 FP16 避免溢出(第 6 章)⭐
充分的 warmup 大 batch 时尤其重要
回滚重跑 尖峰后不恢复 → 回到上一个 checkpoint,换个数据顺序重跑 ⭐
# 🧩 骨架:`model` 来自你自己的代码,这一段只看写法
# ⭐ 跳过异常 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%(空转)但一步都没前进,可能几小时后才发现。

🛑 可以停在这里

⚡ 走神救援

⭐ 千卡训练时硬件故障是常态:单卡平均无故障时间除以卡数,上千张卡就是几小时坏一次——⭐ 所以 checkpoint 和自动恢复决定训练能不能完成。 间隔有个现成的平方根公式,⭐ 它的形状是「保存越贵就存越稀,故障越频就存越密」。

四个要点:⭐ 异步保存(先拷到内存再后台写盘,阻塞从几分钟降到几秒)、⭐ 分片保存(别汇总到一张卡上)、原子写入(先写临时文件再改名,防写一半崩溃)、留多份。

⭐ checkpoint 里最常漏的是数据加载器的位置——⚠️ 漏了会重新从头读数据,某些样本训了好几次、某些一次没见过,而且不报错。 还要存调度器、随机数状态和缩放器。

⭐⭐ 最可怕的不是崩溃,是静默失败——训练照跑、loss 也在降,但结果是错的。⭐ 崩溃是好故障,它至少会告诉你出事了。

⭐ 有个极便宜的探测器:定期把各个 rank 的 loss 和梯度范数收上来比一比——⭐⭐ loss 应该不同(数据不同),但梯度范数在同步之后必须一致;不一致就是通信出了问题。

loss 尖峰的五个手段里性价比最高的是 ⭐ 跳过异常 batch:梯度范数远超近期中位数就丢掉这一步——极便宜、收益很大。⭐ 另一条要记的心态:回滚加换数据顺序重跑,是大厂的标准操作,不是失败。

⚠️⭐ 最后是三个必须设的通信库环境变量:不设阻塞等待的那个,一张卡挂了、其余全部会永远等下去——💀 GPU 显示满负荷,而一步都没走,几小时后才有人发现。

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

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