ZeRO 参数分片:把优化器状态切到 N 片


文档摘要

ZeRO 参数分片:把优化器状态切到 N 片 本节摘要:Adam 对每个参数存两个矩估计,都 float32,一个 70 亿参数模型带 56GB 优化器状态。ZeRO 第一阶段把这切到 N 个 rank,每 rank 持有 1/N 的优化器;本地步进后,更新后的参数分片广播回,每 rank 重建全模型,下一步开始。本节实现 ZeRO-1:用 / 把模型参数打包成连续张量(让按 rank 分片退化为切片),反向后 让每 rank 只收自己分片的梯度、本地 Adam 步进、 回更新参数,并给出 ZeRO-1/2/3 相对朴素 DDP 的显存节省表。

ZeRO 参数分片:把优化器状态切到 N 片

本节摘要: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 章「毕业项目」的分布式训练赛道。

学习目标

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

  1. 把优化器状态(一阶矩、二阶矩、fp32 主副本)切到 N 个 rank,每 rank 持有 1/N。
  2. reduce_scatter 让每 rank 只收自己分片的梯度,再用 allgather 广播更新后参数分片。
  3. 算出 ZeRO 阶段 1、2、3 相对朴素 DDP 的显存节省表
  4. 按模型大小与带宽预算,辩护 ZeRO-1 vs 2 vs 3 的选择。

一、问题与直觉

朴素 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)。显存赢,吞吐持平

二、ZeRO 的阶段

阶段 分片什么 每 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%。

为什么 reduce_scatter 碾压 allreduce-then-shard

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 更新后参数回。
  • 一个 demo:训 3 层 MLP 20 步,逐步打印显存预算对比朴素 DDP 基线。

核心骨架

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 的优化器那一半。
  • 显存数学表:任何 P、N 下算朴素 vs ZeRO-1/2/3 的显存,作为选阶段依据。
  • reduce_scatter + allgather 模式:与 76 节原语对接,证明「净线缆与 DDP 相同,显存被除」。

七、练习

  1. 扩到 ZeRO-2:分片梯度,每 rank 只存自己分片的梯度,反向后把非分片部分置零实现。
  2. 加显存剖析器:在 rank 0 打印实际 fp32 字节用量对比公式预测。
  3. 测墙钟:测朴素 DDP vs ZeRO-1 每步墙钟,分解为 forward、backward、通信。
  4. ZeRO-1 下梯度裁剪:L2 范数必须跨分片用局部范数平方的 allreduce 算。
  5. 朴素 ZeRO:用 allreduce 代替 reduce_scatter,测线缆时间差,用数字辩护 reduce_scatter。

本节要点回顾

  1. ZeRO-1 分片优化器状态:每 rank 持有 1/N 的 fp32 主副本 + Adam 矩。
  2. reduce_scatter 碾压 allreduce-then-shard:后者每 rank 浪费 (N-1)/N;前者恰投递每 rank 的分片。
  3. 净线缆与 DDP 相同:reduce_scatter + allgather 按带宽等于一次 allreduce,显存赢,吞吐持平
  4. 显存数学:朴素 16P,ZeRO-1 是 4P + 12P/N;N=8 降 65%,N=64 降 74%。
  5. 三阶段递进:ZeRO-1(优化器)、ZeRO-2(+梯度)、ZeRO-3(+参数=FSDP,每层通信)。
  6. ZeRO-1 近乎免费:通信同 DDP,显存随 N 线性降,唯一代价簿记;参数显存也成问题才上 2/3。
  7. 混合精度是要点:fp32 主副本才是被分片的,不带混合精度付税不赢。

下一节,我们将进入「流水线并行」——正交的分片轴:不分片优化器状态,而是把层切到 rank,微批流过流水线,核心手艺是压气泡。


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