人人都会AI编程

8.5 训练断点续训、Checkpoint 保存与加载机制

更新时间:2026-07-09

在 8.1 到 8.4 节中,我们梳理了预训练的环境准备、模型初始化、训练执行与监控。但所有这些工作都建立在一个残酷的现实之上:大模型训练极少能一口气跑完。千亿参数模型的预训练周期通常以周甚至月为单位,在此期间,网络抖动、GPU 显存奇偶校验错误(ECC Error)、节点宕机、NCCL 通信超时、存储 I/O 降级等故障不可避免。如果没有可靠的断点续训(Fault Tolerance & Resume)机制,一次宕机就可能导致价值数百万算力投入的前功尽弃。

本节从工程视角拆解 Checkpoint 的完整生命周期:存什么、怎么存、怎么读、以及如何在分布式环境下保证恢复后的训练严格等价于中断前。


一、Checkpoint 的本质:不只是“模型权重”

新手最容易犯的错误,是把 Checkpoint 理解为“保存一下模型的 state_dict”。在大模型预训练场景下,一个可恢复训练状态的 Checkpoint 至少包含以下四类数据

| 组件 | 内容说明 | 体积占比(以 7B 模型为例) |
|------|----------|------------------------|
| Model States | 模型参数(FP16/BF16/FP32)。若使用 LoRA/QLoRA(见 10.2 节),则只存 Adapter 权重 | ~14 GB(BF16) |
| Optimizer States | AdamW 的一阶动量(momentum)和二阶方差(variance),有时还有 Lion、Adafactor 等优化器特有的状态 | ~28–56 GB(FP32 副本) |
| Training States | 全局步数(global step)、当前 epoch、学习率调度器状态、随机数种子(RNG State:PyTorch、CUDA、NumPy 各一份) | 极小(MB 级) |
| Tokenizer & Config | 词表文件、模型配置文件、当前训练超参快照,用于校验一致性 | 极小(MB 级) |

体积估算:使用 AdamW + 混合精度(FP16/BF16 参数 + FP32 主权重 + FP32 动量/方差)时,每 1B 参数的 Checkpoint 总大小约为 4–8 GB。一个 70B 模型的全量 Checkpoint 轻松超过 500 GB,175B 模型则逼近 1.5 TB。这意味着 Checkpoint 的 I/O 本身就是训练流水线中的重大瓶颈。


二、保存策略:同步、异步与分层存储

1. 同步保存的痛点

最简单的实现是在训练脚本里定期调用 torch.save()

if global_step % save_interval == 0:
    torch.save(checkpoint, f"checkpoint-{global_step}.pt")

在同步保存模式下,所有 GPU 会阻塞等待主线程完成写盘。对于大模型,单次保存耗时可能长达数分钟到十几分钟(取决于存储带宽和 Checkpoint 大小),期间算力空转,MFU(算力利用率)暴跌。

2. 异步保存:生产环境的标配

工程上的标准解法是异步 Checkpoint

  • 主训练进程:将状态字典拷贝到主机内存(CPU RAM),立即释放 GPU 继续训练;
  • 后台 I/O 线程/进程:将内存中的 Checkpoint 数据异步写入持久存储(NVMe SSD → 并行文件系统 Lustre/GPFS → 对象存储 S3/OSS)。

主流框架的实现方式:

  • DeepSpeed:通过 deepspeed.save_checkpoint() 接口,支持 ZeRO 分片参数的异步聚合与写入(见下文);
  • PyTorch Distributed Checkpoint (torch.distributed.checkpoint):原生支持异步保存、分片读写,与 FSDP 深度集成;
  • Megatron-LM:通过 --async-checkpointing 参数启用后台线程池写盘。

注意:异步保存并非零成本。CPU 到主存的拷贝仍然需要锁保护,且会短暂占用大量主机内存。若内存不足,可能触发 OOM Killer 反噬训练进程。

3. 保存频率的权衡

| 策略 | 适用场景 | 风险 |
|------|----------|------|
| 高频(每 10–50 步) | 实验初期、小规模调试、不稳定集群 | 存储空间消耗极大,I/O 开销可能拖慢训练 |
| 中频(每 100–1000 步,或每 0.5–2 小时) | 大规模预训练主流做法 | 故障时最多丢失 1–2 小时进度,可接受 |
| 低频(每半天或每天) | 资源极度受限、模型较小 | 一旦故障,损失巨大 |

建议:预训练阶段采用定时触发 + 步数阈值双重兜底(如每 2 小时或每 500 步,以先到者为准)。


三、分布式场景下的 Checkpoint 难题

大模型训练几乎必然使用分布式(详见第 9 章),Checkpoint 的复杂性也随之倍增。

1. 数据并行(DP)与 ZeRO 优化器

  • ZeRO-1/2:优化器状态被分片到不同数据并行 rank 上。保存时,DeepSpeed 默认执行 gather 操作,由 rank 0 聚合完整的优化器状态后落盘。这会导致 rank 0 的内存和存储压力剧增。
  • ZeRO-3:参数、梯度和优化器状态全部分片。保存时必须把所有分片参数聚合回完整模型,才能写入。DeepSpeed 提供了 zero_to_fp32.py 工具用于事后转换,但原生 Checkpoint 体积仍随分片策略膨胀。

工程技巧:若存储空间有限,可配置 DeepSpeed 的 checkpoint_universal_format 或 ZeRO-Infinity,允许各 rank 分别写入自己的分片文件,恢复时再按 rank 加载,避免单点聚合。

2. 张量并行(TP)与流水线并行(PP)

在 Megatron-LM 或类似的 3D 并行框架中:

  • TP:单一层内的权重被横向切分到同节点内的多个 GPU。保存时必须沿张量并行维度 gather 还原完整权重。
  • PP:模型被纵向切成多个 stage,每个 stage 驻留在不同 GPU 上。保存时需要遍历所有 pipeline stage,依次收集各层参数。

这意味着Checkpoint 与并行配置强耦合。一个用 TP=4、PP=8 训练出来的 Checkpoint,不能随意加载到 TP=2、PP=16 的新集群上运行,除非先通过转换脚本重排参数矩阵。


四、恢复训练:热启动与严格续训

Checkpoint 的终极考验不是“能不能存”,而是恢复后训练曲线能否无缝衔接

1. 加载流程的 checklist

从 Checkpoint 恢复时,必须严格校验以下信息,否则可能导致训练发散或静默错误:

  1. 模型结构一致性:层数、隐藏维度、注意力头数、词表大小必须与 Checkpoint 匹配;
  2. 并行配置一致性:TP/PP/DP 分片方式需与保存时一致(或做显式转换);
  3. 优化器状态完整加载:AdamW 的动量/方差若丢失,恢复后第一步更新就会偏离原轨迹;
  4. 随机状态恢复:必须还原 Python、torchtorch.cudanumpy 的 RNG 状态,否则数据读取顺序、Dropout 掩码都会改变,破坏实验可复现性;
  5. 学习率续接:调度器必须从保存时的 step 继续,而不是从头 warmup。

2. 热启动 vs. 冷启动

  • 严格断点续训(Resume):加载完整 Checkpoint(模型+优化器+训练状态),global_step 续接,学习率沿余弦退火曲线继续。这是预训练的标准做法;
  • 热启动微调(Warm-start / Fine-tune):只加载模型权重,优化器和学习率调度器重新初始化。常用于在预训练 Checkpoint 之上做 SFT 或领域适配;
  • 模型转换加载:如将 Megatron-LM 的 Checkpoint 转换为 Hugging Face Transformers 格式用于推理,或反之用于继续训练。

五、生产级最佳实践

1. 存储分级与生命周期管理

Checkpoint 不应只存一份在本地盘:

  • 热层:最近 1–2 个 Checkpoint 保留在训练节点的本地 NVMe SSD,用于秒级故障回滚;
  • 温层:过去 24 小时内的 Checkpoint 推送到集群并行文件系统(Lustre/GPFS);
  • 冷层:每日/每周最优 Checkpoint 归档到对象存储(S3、OSS、HDFS),作为长期备份和后续微调基线。

自动清理:设置保留策略(如只保留最近 5 个 Checkpoint + 每 1000 步的里程碑版本),防止存储爆炸。

2. Checkpoint 校验与心跳监控

  • 保存后校验:写入完成后立即读取文件头,校验 SHA256 或文件大小,避免写入中断产生脏文件;
  • 训练心跳:配合 8.4 节的监控体系,若检测到节点掉线或 NCCL 超时,自动触发最近一次有效 Checkpoint 的恢复脚本,无需人工介入。

3. 小实验 vs. 大生产的差异化策略

| 场景 | 建议策略 |
|------|----------|
| 小规模 SFT / LoRA | 可只保存 LoRA 权重 + 配置,体积仅数百 MB;断点续训需求弱,可接受从头开始 |
| 百亿/千亿基座预训练 | 必须全量异步 Checkpoint;建议每 1–2 小时保存;配置自动故障转移(Failover) |
| 多阶段训练(预训练→SFT→RLHF) | 各阶段结束时保存永久 Checkpoint,并打 tag;阶段内按步数间隔保存临时 Checkpoint |


六、小结

断点续训不是训练流程的“可选项”,而是大模型工程的生命线。其核心在于:

  • Checkpoint 必须包含完整的训练状态(模型+优化器+调度器+随机种子),而不仅是权重;
  • 保存必须是异步的,否则 I/O 阻塞将吞噬宝贵的算力;
  • 分布式 Checkpoint 与并行策略强耦合,恢复时必须严格对齐或做显式转换;
  • 存储需分级管理,平衡快速回滚与长期归档的需求。

在 8.6 节长上下文训练技术之前,掌握 Checkpoint 机制是保障预训练任务“跑得久、摔得起、续得上”的最后一块工程基石。