← 返回题目列表

大模型预训练如何做 Checkpoint 和断点恢复?

中等 第 15 / 25 题 更新于 2026/09/18
大模型预训练Checkpoint分布式训练容错

简化版

可恢复的 Checkpoint 不只保存模型权重,还要保存优化器状态、学习率调度器、梯度缩放器、全局 step、随机数状态、数据游标以及并行切分元数据。写入时采用分片并行、临时目录加完成标记和校验和;恢复后用固定小批次验证 loss 连续、样本不重不漏,再重新加入训练。

详细版

大模型训练一次可能持续数周,节点故障不可避免。Checkpoint 周期是在写盘成本和重算损失间权衡:若故障服从平均间隔 MTBF、一次保存耗时 C,经典近似周期可取 sqrt(2×C×MTBF),再结合对象存储带宽与业务 RPO 调整。

分布式保存有两种路线:聚合成完整权重便于迁移但内存与网络压力大;按 DP/TP/PP rank 分片效率高,但恢复时必须匹配或重分片。恢复流程应先验证 manifest、分片数量与 checksum,再加载模型和优化器,恢复 RNG 与数据采样器,并对保存前后 loss、学习率和样本 ID 做连续性检查。

训练 step -> 各 rank 写临时分片 -> 汇总 manifest/checksum -> 原子发布 COMPLETED
故障恢复 -> 选择最新完整版本 -> 校验 -> 重分片加载 -> 连续性验证 -> 继续训练

完整版教学

一、为什么只保存权重不能续训

权重只描述当前模型参数,而 Adam 还维护一阶矩和二阶矩,学习率调度器保存所处 step,混合精度可能维护 loss scale。缺少这些状态时虽然能“加载模型”,后续更新轨迹却会突变,常出现短暂 loss spike 或收敛质量下降。真正的断点恢复目标是让下一步尽可能等价于未中断训练。

数据状态同样重要。若恢复后从 epoch 开头重新读,部分样本重复;若直接按估算偏移跳过,分布式 shuffle 可能漏样本。应保存 sampler epoch、全局消费 token 数、数据源游标与混合比例状态。

二、完整状态清单

一个可恢复版本通常包含以下内容:

状态缺失后果是否常分片
模型参数/Buffer无法恢复能力
优化器 m、v更新轨迹改变是,且体积大
LR Scheduler、step学习率错位否或小文件
AMP GradScalerFP16 数值波动
CPU/GPU RNGDropout、采样不一致每 rank
数据游标/Sampler样本重复或遗漏每数据 rank
TP/PP/DP 拓扑分片无法解释manifest

还要保存代码版本、Tokenizer、数据清单哈希和超参数。否则半年后即使文件可读,也无法解释它由哪份代码和数据产生。

三、如何决定保存频率

保存过勤会让 GPU 等待 I/O,过稀则故障后重算很多。简化模型中,间隔为 T 时平均丢失约 T/2 的计算,同时每 T 时间支付保存成本 C;Young 公式给出:

T_opt ≈ sqrt(2 × C × MTBF)

若保存需要 4 分钟,集群平均 36 小时发生一次中断,则 T≈sqrt(2×4×2160)≈131 分钟。实际还需考虑恢复耗时、并发故障和存储限流,可先每 2 小时保存,并保留更密集的轻量权重快照。

记忆钩子:Checkpoint 间隔同时在付两种钱——保存时停顿的钱,以及故障后重算的钱;最优点是二者平衡。

四、分片保存与原子发布

ZeRO/FSDP 下每个 rank 只持有参数或优化器的一部分,让所有状态聚合到 rank 0 容易 OOM,也形成网络热点。更可扩展的方式是各 rank 并行写分片,协调者生成 manifest,记录逻辑张量到文件偏移的映射。需要导出部署权重时再离线合并。

保存不能直接覆盖“latest”。各 rank 先写 checkpoint.tmp/<step>,完成后校验大小与 checksum;只有全部分片成功,才写不可变的 COMPLETED 标记并更新 latest 指针。恢复程序只选择有完成标记的版本,从而避开机器中途掉电留下的半成品。

五、改变并行度如何恢复

原来 64 张卡、TP=8 的分片不能假设在 32 张卡、TP=4 下逐文件对应。弹性恢复需要先理解逻辑完整张量,再按新拓扑切分,或使用支持 distributed checkpoint 的格式做 planner 映射。位置编码、词表扩展和 tied weights 还可能有特殊切分规则。

若每个 rank 各自保存私有 pickle 且没有全局元数据,拓扑一变就很难加载。可移植格式应记录张量名、全局 shape、dtype、shard offset 和 replication 关系。重分片要用小模型测试验证拼接维度,避免静默错位。

六、数据与随机性怎样连续

恢复 RNG 包括 Python、NumPy、CPU 和每张 GPU 的随机状态;pipeline 中不同 stage 也可能各有随机流。完全 bitwise 复现还受异步算子和通信顺序影响,但至少应保证 Dropout 与数据抽样不会整体重置。若目标只是统计等价,也要明确可接受范围。

数据集可为每条样本记录稳定 ID,并在 checkpoint 保存最后消费的全局 batch/Token。恢复后抽查边界前后 1000 个 ID,确认无大面积重复或缺口。流式数据若无法 seek,应保存 shard 顺序、文件内 offset,或接受并量化最多重复窗口。

七、恢复演练与保留策略

未经演练的备份不算可靠。应定期在隔离集群从最近版本恢复,跑 10~100 step,并与连续训练对比 loss、梯度范数、学习率和吞吐。还需模拟单分片损坏、latest 指针错误、存储超时和拓扑变化,验证降级到前一版本。

保留可采用“最近 5 个 + 每日 1 个 + 里程碑永久保留”,防止静默坏数据过几天才发现时已无健康版本。删除前确认新版本已异地复制并能读取;权重和训练数据可能敏感,存储需加密和访问审计。

八、常见误区与追问

  • 误区:保存 model.state_dict 就完成了断点恢复。 优化器、调度器、随机数和数据游标缺一都会改变训练轨迹。
  • 误区:文件存在就代表 checkpoint 完整。 分布式写入可能只成功部分 rank,需要 manifest、完成标记和 checksum。
  • 误区:恢复后 loss 没报错就可以继续。 还要检查 step、学习率、样本连续性和梯度范数。
  • 追问:优化器状态为什么特别贵? Adam 的 m、v 常以较高精度保存,体积可达到参数权重的数倍。
  • 追问:如何支持不同 GPU 数恢复? 使用带全局张量元数据的分布式格式,在加载时按新拓扑重分片。
  • 追问:异步保存有什么风险? 内存快照必须对应一致 step,后台写失败要告警,且不能过早复用缓冲区。

九、加强记忆

Checkpoint 可记成“全、分、原、验”:状态要全,模型与优化器要并行分片,临时写完后原子发布,恢复后验证 loss 与数据连续性。再用保存成本和故障间隔估算周期,准备跨拓扑重分片与定期灾备演练,才能把“能加载文件”提升为“训练可以可信地接着跑”。