检查点保存与恢复


文档摘要

检查点保存与恢复 本节摘要:训练中断会杀死运行;检查点让它继续。原子地保存模型、优化器、调度器、损失历史、步计数器与 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)。本节属「预训练/分布式」赛道第六节。

学习目标

阅读完本节,你应当能够:

  1. 把完整训练状态捕获进单一负载,可重载进新进程。
  2. 实现「写临时文件再 rename」的原子保存,使崩溃不留半个文件。
  3. 恢复 Python、NumPy、PyTorch 的 RNG 状态,使恢复后损失匹配无中断基线。
  4. 为不再装进单文件的模型构建分片检查点布局,带哈希校验分片与 JSON 索引。

一、问题与直觉

你设训练 18 小时,墙时间上限 4 小时,集群在 11 小时因某人批的内核升级重启。无检查点你从头来;无恢复你还丢了前 11 小时学的优化器状态——即便模型权重活了,AdamW 矩没了,下一步朝训练轨迹已过的方向踉跄。

五个状态桶

为何重要
模型 权重与缓冲;模型是什么
优化器 动量与自适应矩;无它下一步是不同的优化问题
调度器 学习率在曲线何处;余弦调度尤其在意
训练计数器 步、epoch、epoch 内批,加绘仪表盘的损失历史
RNG 状态 dropout、数据洗牌、模型内任何采样的确定性

原子保存

分片检查点

模型变大时单文件负载太大:加载慢、难检视、网络抖动半读。修复是把参数状态拆分片,写小索引把它们绑一起。

索引记分片数、每分片 sha256、meta 文件 sha256。任一哈希不匹配加载器大声失败。分片可落不同物理盘;meta 小、先读。

恢复续在 epoch 中途

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 Trainersave_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。

五、练习

  1. 分片策略:把轮询分片换成按参数组分片(.weight vs .bias),讨论何时各布局更优。
  2. 保留 K 个:扩保存循环保留最近 K 检查点、剪旧的,盘小时 K 取何值?
  3. 墙时触发:加 --ckpt-every-seconds 按墙时触发保存,而非只步数。
  4. 启动校验:加启动时扫目录每检查点、报哪些损坏的校验路径。
  5. v1→v2 迁移:实现 migrate_v1_to_v2 加新字段、bump schema,使 load 兼容两版本。

本节要点回顾

  1. 五桶:模型/优化器/调度器/训练计数器/RNG,缺一恢复就偏。
  2. 原子保存:写临时文件再 os.replace,rename 在 POSIX 原子。
  3. RNG 是字节:Python/NumPy/torch CPU/torch CUDA 的状态,非只是种子。
  4. epoch 中途续:(epoch, batch_in_epoch) + RNG 快进,损失匹配基线在 1e-4 内。
  5. 分片布局:大模型拆多文件 + meta + JSON 索引(含 sha256),并行读加载。
  6. schema 是升级钩子:字符串在负载里,迁移在它上分支;每分片哈希校验。

下一节,我们做「DDP 与 FSDP」——两个集合通信(broadcast + all-reduce)加一条规则(各 rank 步调一致),从零实现数据并行。


发布者: 作者: Rohit Gupta 转发
评论区 (0)
U