ZeRO 参数分片:把优化器状态切到 N 片 本节摘要:Adam 对每个参数存两个矩估计,都 float32,一个 70 亿参数模型带 56GB 优化器状态。ZeRO 第一阶段把这切到 N 个 rank,每 rank 持有 1/N 的优化器;本地步进后,更新后的参数分片广播回,每 rank 重建全模型,下一步开始。本节实现 ZeRO-1:用 / 把模型参数打包成连续张量(让按 rank 分片退化为切片),反向后 让每 rank 只收自己分片的梯度、本地 Adam 步进、 回更新参数,并给出 ZeRO-1/2/3 相对朴素 DDP 的显存节省表。
本节摘要:Adam 对每个参数存两个矩估计,都 float32,一个 70 亿参数模型带 56GB 优化器状态。ZeRO 第一阶段把这切到 N 个 rank,每 rank 持有 1/N 的优化器;本地步进后,更新后的参数分片广播回,每 rank 重建全模型,下一步开始。本节实现 ZeRO-1:用
flatten_params/unflatten_into把模型参数打包成连续张量(让按 rank 分片退化为切片),反向后reduce_scatter让每 rank 只收自己分片的梯度、本地 Adam 步进、allgather回更新参数,并给出 ZeRO-1/2/3 相对朴素 DDP 的显存节省表。读完本节,你能说清为什么 reduce_scatter 碾压 allreduce-then-shard(后者每 rank 浪费 (N-1)/N 的归约)、为什么 ZeRO-1 是近乎免费的赢(通信按带宽与 DDP 相同),以及按模型大小与带宽预算在三个阶段间选择的依据。
对应原课程:Phase 19 · Lesson 78 ·
zero-parameter-sharding(原英文phases/19-capstone-projects/78-zero-parameter-sharding/docs/en.md)。本节属第 20 章「毕业项目」的分布式训练赛道。
阅读完本节,你应当能够:
朴素 DDP 复制一切:参数、梯度、优化器状态在每 rank 全量存在。一个 70 亿参数模型用 fp16,即每 rank 14GB 参数 + 14GB 梯度 + 28GB 优化器状态。优化器状态是最大项,也最易分片,因为它只在步进时碰,不在 forward 或 backward 碰。
ZeRO-1 分片优化器状态,每 rank 持有 1/N 的 Adam 矩。反向后,不 allreduce 全梯度再本地步进,而是 reduce_scatter 让每 rank 只收自己分片的求和梯度;该 rank 对自己分片的主参数跑优化器步进;更新后参数分片再 allgather 回,让每 rank 拥有下一步 forward 的全模型。优化器显存降 N 倍,每步线缆流量与 DDP 相同(reduce_scatter 加 allgather 按带宽等于一次 allreduce)。显存赢,吞吐持平。
| 阶段 | 分片什么 | 每 rank 显存 | 每步通信 |
|---|---|---|---|
| DDP | 无 | params + grads + optim | 1x allreduce |
| ZeRO-1 | 优化器状态 | params + grads + optim/N | 1x reduce_scatter + 1x allgather |
| ZeRO-2 | 优化器 + 梯度 | params + grads/N + optim/N | 1x reduce_scatter + 1x allgather |
| ZeRO-3 | 优化器 + 梯度 + 参数 | params/N + grads/N + optim/N | 每层 1x allgather + 每层 1x reduce_scatter |
阶段 1 是最便宜的赢,因为优化器状态主导预算;阶段 2 需梯度分片累加逻辑但带宽相同;阶段 3(FSDP)为每 forward 与 backward 付每层通信,换取参数分片的显存下降。本节完整实现阶段 1。
对 P 个参数、Adam 混合精度训练:
| 项 | 朴素 | ZeRO-1 | 为什么 |
|---|---|---|---|
| fp16 参数 | 2P 字节 | 2P 字节 | forward 需要 |
| fp16 梯度 | 2P 字节 | 2P 字节 | backward 需要 |
| fp32 主副本 | 4P 字节 | 4P/N 字节 | 只有优化器用 |
| fp32 一阶矩 | 4P 字节 | 4P/N 字节 | 只有优化器用 |
| fp32 二阶矩 | 4P 字节 | 4P/N 字节 | 只有优化器用 |
| 合计 | 16P 字节 | 4P + 12P/N 字节 |
N=8:朴素 16P,ZeRO-1 5.5P,降 65%;N=64:朴素 16P,ZeRO-1 4.19P,降 74%。
allreduce 给每 rank 完整求和梯度。若你只需要分片 r,被归约的 (N-1)/N 在 rank r 上浪费了。reduce_scatter 恰投递每 rank 拥有的分片,每 rank 字节与 allreduce 相同(因 allreduce = reduce_scatter + allgather),但后半被稍后的参数分片 allgather 替换。净线缆与 DDP 相同,显存被除。
code/main.py 实现:
flatten_params(module) 与 unflatten_into(module, flat):把模型参数打包成一个连续张量并解包回。扁平布局让按 rank 分片退化为简单切片。ZeroOptimizer(model, world_size, rank, lr):持有该 rank 的主副本与 Adam 矩分片。step():对扁平梯度 reduce_scatter,对该 rank 分片跑 Adam,allgather 更新后参数回。class ZeroOptimizer: def __init__(self, model, world_size, rank, lr): self.flat = flatten_params(model) # 连续扁平张量 self.shard = self.flat.size(0) // world_size # 每 rank 只持有自己分片的 fp32 主副本与 Adam 矩 self.master = self.flat.float()[rank*self.shard:(rank+1)*self.shard].clone() self.m = torch.zeros_like(self.master) self.v = torch.zeros_like(self.master) def step(self, flat_grad): shard_grad = reduce_scatter(flat_grad, self.world_size) # 只收自己分片 # 对自己分片跑 Adam self.m = 0.9*self.m + 0.1*shard_grad self.v = 0.999*self.v + 0.001*shard_grad**2 self.master -= lr * self.m / (self.v.sqrt() + 1e-8) # allgather 更新后参数,让每 rank 重建全模型 updated_flat = allgather(self.master, self.world_size) unflatten_into(self.model, updated_flat)
运行 python3 code/main.py,输出逐步 loss 与显存表,显示 ZeRO-1 每 rank 持有 1/N 优化器状态对比 DDP 的全副本。
业界对比:DeepSpeed ZeRO 是参考实现,deepspeed_config.json 选阶段 1/2/3 与分区大小;PyTorch FSDP 是 PyTorch 原生等价物,ShardingStrategy.SHARD_GRAD_OP 即 ZeRO-2,FULL_SHARD 即 ZeRO-3;HuggingFace Accelerate 在统一配置下包两者。它们的共识与本节一致:阶段 1 是近乎免费的赢(通信按带宽与 DDP 相同,显存随 N 线性降,唯一代价是优化器分片的簿记),除非参数分片显存也成问题才上阶段 2/3 用通信换显存。
三种模式把 ZeRO 硬化到可上线。
分片检查点要紧。 ZeRO-1 的优化器状态跨 rank 切分,检查点必须记录哪个 rank 拥有什么。80 节构建分片检查点清单,在同 world size 上恢复 ZeRO 运行;没有它,保存的状态在重启时不可读。
混合精度是要点。 ZeRO 是混合精度技术,fp32 主副本才是被分片的。不带混合精度跑 ZeRO,fp32 主副本的显存税付了却没拿到对应的 fp16 forward 赢。生产运行总是把 ZeRO 与 autocast 或 bf16 权重配对。
阶段 1 是近乎免费的赢。 通信按带宽与 DDP 相同,显存随 N 线性降,唯一代价是簿记。生产栈默认阶段 1,除非参数分片显存也成问题;那时阶段 2 或 3 用通信换显存。
ZeroOptimizer + flatten_params/unflatten_into:本节是 80 节分片检查点的保存对象、81 节端到端 demo 的优化器那一半。下一节,我们将进入「流水线并行」——正交的分片轴:不分片优化器状态,而是把层切到 rank,微批流过流水线,核心手艺是压气泡。