分片检查点与原子恢复:让万步训练不怕中断


文档摘要

分片检查点与原子恢复:让万步训练不怕中断 本节摘要:一个 700 亿参数训练任务每隔几小时就被节点故障打断,检查点格式决定你丢 30 分钟还是 30 小时。分片检查点让每 rank 并行写自己的分片、在清单里记录归属;恢复时每 rank 从自己的文件加载分片,在同 world size 上重建状态,优化器像什么都没发生一样步进;原子写(写临时路径再 rename)防止半成品检查点毒化下次恢复。本节实现 清单(worldsize、sha256、paramshardoffset/numel)、原子 temp-then-rename 写、加载时 sha256 校验与往返字节等价测试,并辩护清单 schema 对三种失败模式(world size 变、分片数不匹配、部分写)的防御。

分片检查点与原子恢复:让万步训练不怕中断

本节摘要:一个 700 亿参数训练任务每隔几小时就被节点故障打断,检查点格式决定你丢 30 分钟还是 30 小时。分片检查点让每 rank 并行写自己的分片、在清单里记录归属;恢复时每 rank 从自己的文件加载分片,在同 world size 上重建状态,优化器像什么都没发生一样步进;原子写(写临时路径再 rename)防止半成品检查点毒化下次恢复。本节实现 ShardManifest 清单(world_size、sha256、param_shard_offset/numel)、原子 temp-then-rename 写、加载时 sha256 校验与往返字节等价测试,并辩护清单 schema 对三种失败模式(world size 变、分片数不匹配、部分写)的防御。读完本节,你能说清为什么单文件 gather-then-write 让 1TB 检查点过一 rank 网口要 4 小时而分片 64 rank 只要 4 分钟。

对应原课程:Phase 19 · Lesson 80 · checkpoint-sharded-resume(原英文 phases/19-capstone-projects/80-checkpoint-sharded-resume/docs/en.md)。本节属第 20 章「毕业项目」的分布式训练赛道。

学习目标

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

  1. 把多 rank 检查点存成每 rank 一个分片文件加一份记录归属的清单
  2. 原子写模式(写临时路径再 rename),使写中途崩溃永不产出半成品检查点。
  3. 从清单恢复,在每 rank 上对 fp16 参数与 ZeRO 优化器状态验证字节等价
  4. 辩护清单 schema 对三种失败模式的防御:world size 变、分片数不匹配、部分写

一、问题与直觉

朴素检查点把所有参数与优化器状态读进 rank 0、gather、写单文件。对 700 亿模型那是 1.1TB 状态过一个 rank 的网口;写阻塞其他所有 rank(它们闲等 gather),IO 带宽是最慢单 GPU 的网路链路而非聚合。真实集群上 gather-then-write 能比前一个训练小时还久,意味着任务每天不到一个检查点。

分片检查点翻转模式:每 rank 并行写自己的分片到自己的文件,清单记录哪个 rank 拥有哪个分片以便恢复时各归各位。聚合写带宽随集群缩放——过一 rank 4 小时的 1TB 检查点,过 64 rank 只要 4 分钟。加清单给你不兼容恢复的契约:world size 变可测、部分写可测,加载路径能大声失败而非静默用陈旧数据。

二、从零实现

清单 schema

{ "world_size": 4, "step": 1234, "wall_clock_seconds": 4521, "shards": [ {"rank": 0, "path": "rank0.bin", "sha256": "...", "param_shard_offset": 0, "param_shard_numel": 65536}, {"rank": 1, "path": "rank1.bin", "sha256": "...", "param_shard_offset": 65536, "param_shard_numel": 65536} ], "schema_version": 1 }

三个字段承重:world_size 让不同规模的恢复大声失败而非静默损坏;每分片 sha256 抓部分或损坏写;每分片 param_shard_offsetparam_shard_numel 让加载器在正确位置重建扁平参数张量。

原子写

标准模式:每分片写到 <name>.tmp,清单写到 manifest.json.tmp,各 fsync,再 rename。POSIX rename 在同一文件系统内是原子的——新文件要么全在要么旧文件还在。最终 rename 前崩溃,前一个检查点仍是活的。没有原子写,崩溃可留下部分分片加指向它的清单,恢复时优化器状态损坏。

清单须防御的三种失败

失败 症状 防御
World size 变 N=8 上恢复 N=4 的清单 清单里 world_size 不匹配,大声失败
分片数不匹配 恢复见 rank*.bin 少于清单分片 枚举分片,验证每个存在
部分写 分片文件刷新中途截断 加载时 sha256 校验

每种防御早拒坏加载;替代是 100 步后 loss 变 NaN 才浮现的静默损坏。

为什么每 rank 一文件而非一个大文件

通过 O_APPEND 并发写一文件在 POSIX 上对字节对齐写可行,但实践上单分片内偏移跨 MB 级区域,锁主导。每 rank 一文件无争用,并在底层文件系统并行时(Lustre、GPFS)受益于条带化。生产栈(DeepSpeed、FSDP、NeMo)全用每 rank 文件。

实现骨架

@dataclass class ShardManifest: world_size: int; step: int; shards: list; schema_version: int = 1 def save_sharded(state_per_rank, directory, step, world_size): for rank, state in enumerate(state_per_rank): _atomic_write(directory / f"rank{rank}.bin", state) # tmp + rename manifest = ShardManifest(world_size, step, [ {"rank": r, "path": f"rank{r}.bin", "sha256": _sha256(state_per_rank[r]), "param_shard_offset": r*shard_numel, "param_shard_numel": shard_numel} for r in range(world_size)]) _atomic_write(directory / "manifest.json", manifest.to_json()) def load_sharded(directory, expected_world_size): manifest = ShardManifest.from_json(read(directory / "manifest.json")) if manifest.world_size != expected_world_size: raise ValueError(f"world size 不匹配:清单 {manifest.world_size} ≠ 期望 {expected_world_size}") states = [] for shard in manifest.shards: if not (directory / shard["path"]).exists(): raise FileNotFoundError(f"缺分片 {shard['path']}") data = read(directory / shard["path"]) if _sha256(data) != shard["sha256"]: raise ValueError(f"分片 {shard['path']} sha256 校验失败") states.append(data) return states

code/main.py 实现:ShardManifest dataclass(含 to_json/from_json)、save_sharded(原子 temp-then-rename 写每 rank 二进制)、load_sharded(读清单、校验每分片 sha256、返回每 rank 状态字典)、往返测试(建状态、存、加载、断言字节等价)。运行 python3 code/main.py 输出 4 个分片文件加清单,再字节等价重载。

三、框架对比

业界对比:DeepSpeed checkpointingdeepspeed.save_checkpoint(tag=step) 写每 rank 文件加指向活动 tag 的 latest 文件;PyTorch FSDPtorch.distributed.checkpoint 用决定每 rank 布局的 Planner 存分片状态;NeMo 包 DeepSpeed 与 FSDP,统一 save_to_checkpoint API 加元数据。它们的共识与本节一致:每 rank 文件、清单记归属、原子写、sha256 校验、world size 变大声失败

四、生产中的硬化模式

三种模式把检查点硬化到可上线。

异步写。 生产栈在独立线程或进程上发检查点写,训练继续。屏障在下次检查点:上次没完别开下次。DeepSpeed 的 async_io 标志正是此意;本节保持同步以让步骤可见。

本地快盘先,再异步上传。 先写本地 NVMe(快)再异步上传 S3 或 GCS。两层模式让集群内检查点恢复快、同时把耐久副本送出集群归档。清单带本地路径,上传清单带远程路径。

轮转要紧。 生产运行留最近 K 个检查点(常 3~5),轮转最旧的。不轮转磁盘中途满,下次检查点失败;轮转则下次存先删最旧,腾预算。

五、可复用产物

  • ShardManifest + save_sharded/load_sharded:本节是 78 节 ZeRO 状态与 81 节端到端 demo 的存取形状;清单的 world_size/sha256/offset 三字段是任何分片检查点的最小契约。
  • 原子 temp-then-rename 模式:可独立用于任何「崩溃不留半成品」的写场景。
  • 三失败防御表:作为检查点设计的审计清单。

六、练习

  1. 加异步写:在线程里 kick off 存盘,训练继续,下次存盘阻塞至上次完成。
  2. last_5_steps 轮转:留 5 个最近检查点,存新的前删最旧。
  3. 加 CRC 快校验:内循环重载(轮转把检查点转为活动)用 CRC-only 而非全 sha256。
  4. 跨 world size 加载:从 N=4 分片再分片到 N=8,读清单、拼接、再切。
  5. 加上传假 S3(第二目录)写上传清单,辩护两层存储策略。

本节要点回顾

  1. 检查点格式决定丢 30 分钟还是 30 小时——节点故障每几小时打断一次,万步训练必须可恢复。
  2. 分片检查点翻转 gather-then-write:每 rank 并行写自己分片,聚合带宽随集群缩放(1TB 过一 rank 4 小时 → 64 rank 4 分钟)。
  3. 清单三承重字段:world_size(变大声失败)、sha256(抓部分写)、param_shard_offset/numel(重建扁平张量)。
  4. 原子 temp-then-rename:POSIX rename 同文件系统内原子,崩溃留前一个为活检查点。
  5. 三失败防御:world size 变、分片数不匹配、部分写——每种早拒,替代是 100 步后 NaN 的静默损坏。
  6. 每 rank 一文件:无锁争用、并行文件系统受益条带化,生产栈(DeepSpeed/FSDP/NeMo)共识。
  7. 三大硬化:异步写(独立线程)、两层存储(本地 NVMe + 异步云)、轮转(留 K 删最旧)。

下一节,我们将进入「端到端分布式训练」——把 76~80 节组装成一个跑在 4 模拟 rank 上的微型 GPT:DDP 同步梯度、ZeRO-1 分片优化器、中点分片检查点,20 步自终止。


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