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