检查点保存与恢复 本节摘要:训练中断会杀死运行;检查点让它继续。原子地保存模型、优化器、调度器、损失历史、步计数器与 RNG 状态,使任意时刻的崩溃在磁盘留下有效文件。右产物是单个文件,持有继续所需的一切:模型参数、优化器状态、调度器状态、绘图用的损失历史、当前步/epoch/epoch 内批计数器、与每个随机源的 RNG 状态。无 RNG 状态,恢复后的损失曲线是另一条曲线——同模型同数据、不同洗牌、不同 dropout 掩码、仪表盘上不同数字。原子保存是契约另一半:写进临时文件再 rename,崩溃留前一个好文件;POSIX 上 rename 原子。 对应原课程:Phase 19 · Lesson 47 · (原英文 )。本节属「预训练/分布式」赛道第六节。
本节摘要:训练中断会杀死运行;检查点让它继续。原子地保存模型、优化器、调度器、损失历史、步计数器与 RNG 状态,使任意时刻的崩溃在磁盘留下有效文件。右产物是单个文件,持有继续所需的一切:模型参数、优化器状态、调度器状态、绘图用的损失历史、当前步/epoch/epoch 内批计数器、与每个随机源的 RNG 状态。无 RNG 状态,恢复后的损失曲线是另一条曲线——同模型同数据、不同洗牌、不同 dropout 掩码、仪表盘上不同数字。原子保存是契约另一半:写进临时文件再 rename,崩溃留前一个好文件;POSIX 上 rename 原子。
对应原课程:Phase 19 · Lesson 47 ·
checkpoint-save-resume(原英文phases/19-capstone-projects/47-checkpoint-save-resume/docs/en.md)。本节属「预训练/分布式」赛道第六节。
阅读完本节,你应当能够:
你设训练 18 小时,墙时间上限 4 小时,集群在 11 小时因某人批的内核升级重启。无检查点你从头来;无恢复你还丢了前 11 小时学的优化器状态——即便模型权重活了,AdamW 矩没了,下一步朝训练轨迹已过的方向踉跄。
| 桶 | 为何重要 |
|---|---|
| 模型 | 权重与缓冲;模型是什么 |
| 优化器 | 动量与自适应矩;无它下一步是不同的优化问题 |
| 调度器 | 学习率在曲线何处;余弦调度尤其在意 |
| 训练计数器 | 步、epoch、epoch 内批,加绘仪表盘的损失历史 |
| RNG 状态 | dropout、数据洗牌、模型内任何采样的确定性 |
模型变大时单文件负载太大:加载慢、难检视、网络抖动半读。修复是把参数状态拆分片,写小索引把它们绑一起。
索引记分片数、每分片 sha256、meta 文件 sha256。任一哈希不匹配加载器大声失败。分片可落不同物理盘;meta 小、先读。
snap 到下一 epoch 起点的恢复浪费几分到一天。修复是 (epoch, batch_in_epoch) 加 RNG 状态。加载后训练循环快进 RNG 跳过当前 epoch 已消费的批,从 batch_in_epoch 继续。断言:恢复后损失轨迹匹配无中断基线在 1e-4 内。
code/main.py 提供四个原语加 demo 驱动。
捕获与恢复 RNG:capture_rng_state 返回含 Python random.getstate、NumPy np.random.get_state、PyTorch CPU 与 CUDA RNG 字节的字典;restore_rng_state 反转。CPU 张量是 PyTorch RNG 知道如何消费的 uint8 字节缓冲。
原子保存:atomic_save 写负载到目标目录的临时文件,os.replace 换进终名;atomic_write_json 对分片索引同理。
全往返:
def save_checkpoint(path, model, opt, sched, train_state, schema="v1"): payload = { "schema": schema, # 升级钩子 "model": model.state_dict(), "optimizer": opt.state_dict(), "scheduler": sched.state_dict(), "train_state": asdict(train_state), # step/epoch/batch/losses "rng": capture_rng_state(), "wall_saved_at": time.time(), } atomic_save(payload, path) def load_checkpoint(path, model, opt, sched): payload = torch.load(path, map_location="cpu") if payload["schema"] != "v1": # 版本分发 raise ValueError(f"unsupported schema {payload['schema']}") model.load_state_dict(payload["model"]) opt.load_state_dict(payload["optimizer"]) sched.load_state_dict(payload["scheduler"]) restore_rng_state(payload["rng"]) return TrainState(**payload["train_state"])
分片变体:save_sharded_checkpoint 把参数键轮询分到 N 片,每片独立原子保存,写 meta(优化器/调度器/训练态),写 JSON 索引(分片 sha256);load_sharded_checkpoint 合并前校验每片。
恢复 demo:run_resume_demo 训练 total_steps,在 interrupt_at 存检查点,再继续;第二进程恢复检查点跑剩余步,返回两损失轨迹在中断点后的最大绝对差——RNG 恢复时为零或浮点噪声。
设计要点:
schema是字符串在负载里,迁移在它上分支,无它你无法演化格式而不破旧运行。每分片 sha256——静默截断下载是最坏的 bug,加载器要么快失败要么晚失败。检查点节拍诚实:每 N 步且每墙分钟存一次(取短),否则崩的长步浪费整窗工作。
HuggingFace Trainer 的 save_strategy="steps"、PyTorch Lightning 的 ModelCheckpoint、Megatron-LM/DeepSpeed 的分布式检查点都是同形状:模型+优化器+调度器+计数器+RNG,原子写,按步命名使最新易找。分片布局驱动大模型并行读加载;index.json 是让它工作的件。本节手写让你看清五桶、原子 rename、RNG 字节捕获、分片轮询。生产栈还处理:跨设备 map_location、BF16 检查点、FSDP 分片对齐(分片边界须匹配拓扑),但核心契约与本节一致。
code/main.py + outputs/skill-checkpoint-save-resume.md(任何新训练脚本的食谱:负载形状、原子写、RNG 捕获、分片索引)。demo 单文件与分片都断言 max-diff < 1e-4,摘要落 outputs/resume-demo.json。把 save_checkpoint 接周期保存点、load_checkpoint 接启动,运行就扛得住 kill。
.weight vs .bias),讨论何时各布局更优。--ckpt-every-seconds 按墙时触发保存,而非只步数。migrate_v1_to_v2 加新字段、bump schema,使 load 兼容两版本。os.replace,rename 在 POSIX 原子。(epoch, batch_in_epoch) + RNG 快进,损失匹配基线在 1e-4 内。下一节,我们做「DDP 与 FSDP」——两个集合通信(broadcast + all-reduce)加一条规则(各 rank 步调一致),从零实现数据并行。