端到端分布式训练:DDP + ZeRO + 分片检查点的组装 本节摘要:76 到 80 节各建一片,本节是组装:一个跨 4 个模拟 rank 训练的微型 GPT,用 DDP 同步梯度、ZeRO-1 分片优化器状态、中点存分片检查点。demo 跑 20 步,自终止,打印 loss 曲线加显存画像,写出可恢复检查点,并验证四个不变量——loss 在浮点噪声内单调下降、每 rank 每步参数范数相同、每 rank 优化器显存等于 ZeRO-1 公式 12P/N 字节、第 10 步检查点重载字节等价。读完本节,你交付的不是一个 demo,而是一个分布式训练子系统:每片在前面课独立可测,本节证明它们组装后仍正确,这正是真实团队采纳 DeepSpeed 前要自建的那套抽象。
本节摘要:76 到 80 节各建一片,本节是组装:一个跨 4 个模拟 rank 训练的微型 GPT,用 DDP 同步梯度、ZeRO-1 分片优化器状态、中点存分片检查点。demo 跑 20 步,自终止,打印 loss 曲线加显存画像,写出可恢复检查点,并验证四个不变量——loss 在浮点噪声内单调下降、每 rank 每步参数范数相同、每 rank 优化器显存等于 ZeRO-1 公式 12P/N 字节、第 10 步检查点重载字节等价。读完本节,你交付的不是一个 demo,而是一个分布式训练子系统:每片在前面课独立可测,本节证明它们组装后仍正确,这正是真实团队采纳 DeepSpeed 前要自建的那套抽象。
对应原课程:Phase 19 · Lesson 81 ·
end-to-end-distributed-train(原英文phases/19-capstone-projects/81-end-to-end-distributed-train/docs/en.md)。本节属第 20 章「毕业项目」的分布式训练赛道,是本赛道的收官。
阅读完本节,你应当能够:
毕业项目是「部件拼得起来」的证明。76 节实现集体通信,77 节包成 DDP,78 节用 reduce_scatter 分片优化器状态,79 节分析流水线,80 节存分片检查点——每节独立、各有测试。真实训练同时用每个原语;组装错了,要么 loss 发散,要么检查点拒绝恢复,要么每 rank 显存该缩反涨。
本节跑端到端 demo 并验证四个不变量:(a) 20 步内 loss 在浮点噪声内单调下降;(b) 每步每 rank 持相同参数范数;(c) 每 rank 优化器显存等于 ZeRO-1 公式 12P/N 字节;(d) 第 10 步检查点重启字节等价重载。demo 自终止:20 步、单命令、退出码 0。
模型故意小:2 个 transformer 块、嵌入维 32、4 个注意力头、词表 64、序列长 16、batch 4,几千参数。大到跑遍每个接线决策(多头注意力跑标准掩码路径、LayerNorm 有权重要同步、LM 头是回词表的独立线性投影),小到 4 CPU rank 上 20 步秒级完成。
| 课的片 | 拥有什么 | 留给回路什么 |
|---|---|---|
| DDP 广播 | 初始参数同步 | 构造时一次调用 |
| ZeRO-1 步 | 梯度同步、主副本更新、参数广播 | 每步一次调用,替换 optimizer.step |
| 分片检查点 | 持久化每 rank 状态、带 sha256 清单 | 在 rank 0 上调用,状态经 allgather 收集 |
| 训练回路 | forward、backward、loss 记录 | 按序调上面三者 |
回路不知道 reduce_scatter 或 rendezvous 文件——ZeRO 与检查点模块暴露窄接口供回路组装。
77 节的 MLP 足以验证梯度同步。微型 GPT 加三样东西:词表上的独立 LM 头(本节为清晰不 tying,完整 GPT 通常把头与 token 嵌入 tying)、softmax+交叉熵损失(比 MSE 多数值边界)、非对称 forward(嵌入→注意力→每层 MLP)。毕业项目继续用 MLP 会藏掉组装是否正确处理 LayerNorm 或嵌入层的梯度形状。
def _train_worker(rank, world_size, ...): init_process(rank, world_size, backend="gloo") torch.manual_seed(seed) # 参数 init 共享种子 model = MiniGPT() broadcast_params(model, src=0) # DDP 广播 torch.manual_seed(seed + rank) # shuffle rank 专属种子 zero_opt = ZeroOptimizer(model, world_size, rank, lr) for step in range(20): loss = model(batch_shard(rank)).cross_entropy() loss.backward() zero_op.step(flat_grad(model)) # ZeRO-1:reduce_scatter + Adam + allgather if step == 10 and rank == 0: save_sharded(collect_state(model, zero_op), dir, step, world_size) if step == 10: verify_resume(dir, world_size, zero_op) # 字节等价
code/main.py 实现:MiniGPT(2 层 transformer,带掩码自注意力与独立 LM 头)、make_corpus(确定性下一 token 预测数据)、_train_worker(每 rank spawn;广播 init 参数、跑回路、调 ZeRO 步、第 10 步写分片检查点)、verify_resume(主运行后在进程内重载第 10 步检查点,断言保存的主分片与内存快照逐字节匹配)、main(编排全 demo,打印 loss 表、显存画像、验证结果)。运行 python3 code/main.py 输出 20 行 loss 表、4 行每 rank 显存画像、检查点清单、成功的「RESUME VERIFIED」行。
💡 自终止意味着退出 0:回路跑固定 20 步退出,无
while True、无人干预、无外部状态恢复。一个你能放任自流跑完拿到完整日志的毕业项目,才是证明系统接线正确的毕业项目——任一片死锁,demo 永不返回,测试架抓到。
业界对比:DeepSpeed 在一个配置下组合 DDP + ZeRO + 流水线 + 激活检查点,本节的组装是 DeepSpeed 形状的微缩;PyTorch FSDP 是原生等价物,FullyShardedDataParallel 配 ShardingStrategy.SHARD_GRAD_OP 即 ZeRO-2;NeMo 与 Megatron-LM 对最大模型加张量并行,否则组装同形。它们的共识与本节一致:毕业项目是「部件拼得起来」的证明,每片独立可测,组装后验证不变量(loss/参数范数/显存公式/检查点恢复)。
三种模式为真实运行完成组装。
每 K 分钟而非每 K 步存检查点。 步时间随序列长度与微批数变。10 分钟检查点节拍不论模型大小抓同样计算。本节为简用步基;生产用墙钟基。
早检测发散。 生产运行在反向后加 NaN 守卫与 loss 飙升检测器;loss 一步跳超 2 倍就回滚到前一检查点,而非让优化器走入退化态。本节 loss 曲线平滑故守卫闲置,但钩子留着。
跨 rank 聚合显存画像。 真实运行每 rank 显存不同(最大流水线段的 rank 持更多激活)。生产记跨 rank 的 max 加 mean;本节打印每 rank 以示公式匹配。
MiniGPT + _train_worker + verify_resume:本节是分布式训练赛道的集成点;76~80 节每节拥有一个模块,本节导入而非复制。这是真实团队采纳 DeepSpeed 前要自建的那套抽象的微缩。6 节课合起来就是真实团队在采纳 DeepSpeed 前会自建的分布式训练子系统;抽象已在 gloo 上验证,失败模式已演练。第 17 章(基础设施与生产)是把它搬到真实集群的地方。
下一节,我们将切换到安全赛道(Track I),从「越狱分类学」开始——把对抗性攻击按手法分类(角色扮演、越权、编码、多轮),为后续检测器、拒绝评估、安全门打地基。