在 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 恢复时,必须严格校验以下信息,否则可能导致训练发散或静默错误:
- 模型结构一致性:层数、隐藏维度、注意力头数、词表大小必须与 Checkpoint 匹配;
- 并行配置一致性:TP/PP/DP 分片方式需与保存时一致(或做显式转换);
- 优化器状态完整加载:AdamW 的动量/方差若丢失,恢复后第一步更新就会偏离原轨迹;
- 随机状态恢复:必须还原 Python、
torch、torch.cuda、numpy的 RNG 状态,否则数据读取顺序、Dropout 掩码都会改变,破坏实验可复现性; - 学习率续接:调度器必须从保存时的 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 机制是保障预训练任务“跑得久、摔得起、续得上”的最后一块工程基石。