1910.02054-zero-memory-optimizations-trillion-parameter-models

ZeRO: Memory Optimizations Toward Training Trillion Parameter Models

ZeRO 的核心贡献是把数据并行中每张 GPU 都完整复制的 optimizer states、gradients、parameters 拆成可按 data-parallel rank 分片保存和按需通信的动态状态系统,使 data parallelism 获得接近 model parallelism 的内存效率,同时保留 data parallelism 的大粒度计算和接近原始 DP 的通信量;论文实际评估了分片 optimizer states 与 gradients、再配合 ZeRO-R 的 ZeRO-100B,实现 400 张 V100 上训练 100B+ GPT-like 模型、最高 170B 可运行、100B 规模约 15 PFLOPs 吞吐,并把进一步分片 parameters 的完整方案作为通向 trillion-parameter training 的内存与通信分析。

Authors Samyam Rajbhandari, Jeff Rasley, Olatunji Ruwase, Yuxiong He

已审阅 Archived 2026-06-18 22:13 Updated 2026-07-16 19:17 Reviewed 2026-07-18 17:42 Source

Source

作者与关系

阅读目标与判断边界

本笔记关注:

  1. ZeRO 如何重新定义大模型训练中的“显存瓶颈”,尤其是 optimizer state、gradient、parameter 三类 model states 的冗余。
  2. ZeRO-DP 三阶段分片和 ZeRO-R residual memory optimizations 的具体计算、通信和工程权衡。
  3. 它和当前本地档案里 scaling laws、FlashAttention、RLHF/RLVR systems、slime/VERL、GLM/DeepSeek 系列报告的关系。

判断边界:

  • 本笔记分析的是 arXiv v3 论文。DeepSpeed 后续版本里的 ZeRO-Offload、ZeRO-Infinity、ZeRO-3 implementation details、FSDP/DTensor 等演进只作为背景关系,具体实现不归入本文作者原始结果。
  • 论文的 1T 结论主要来自内存和通信模型;实测重点是 ZeRO-100B,即 Pos+gP_{os+g} 加 ZeRO-R。
  • 论文实验基于 400 张 V100、DGX-2、Megatron-LM 2019 版本和 GPT-like dense Transformer;当前 H100/B200、MoE、FSDP、Megatron-Core、DeepSpeed 后续版本会改变绝对吞吐和 baseline。

论文脉络

1. 研究问题、背景和价值

2019 年前后,NLP 模型规模从 BERT-Large、GPT-2 推进到 Megatron-LM 8.3B、T5 11B。模型变大带来精度收益,但训练系统遇到一个更基础的约束:单卡显存存不下训练所需状态。直觉上,1.5B 参数的 GPT-2 只需要约 3GB 的 fp16 参数空间;实际训练时,一张 32GB V100 仍然很容易 OOM。论文要回答的问题是:显存到底被哪些状态占掉了,能否让更多 GPU 的 aggregate memory 真正变成单个大模型可用的训练内存。

这个问题的重要性来自两层价值。第一层是能力 scaling:如果模型 loss 和能力随参数、数据、计算平滑改善,那么训练系统必须让研究者能实际运行更大模型。第二层是工程可用性:model parallelism 和 pipeline parallelism 可以降低单卡显存,但需要改模型、引入跨层/跨算子通信、改变 batch 或产生 pipeline bubble;普通 data parallelism 易用、高效,却在每张 GPU 上完整复制所有 model states。ZeRO 尝试保留 DP 的粗粒度计算和易用性,同时把重复状态切分掉。

论文把训练显存拆成两大类:

  1. Model states:parameters、gradients、optimizer states。
  2. Residual states:activations、temporary buffers、fragmented memory。

其中 model states 是本文的主战场。以 mixed-precision Adam 为例,Ψ\Psi 个参数需要:

2Ψfp16 param+2Ψfp16 grad+4Ψfp32 master+4Ψmomentum+4Ψvariance=16Ψ bytes. 2\Psi_{\mathrm{fp16\ param}} + 2\Psi_{\mathrm{fp16\ grad}} + 4\Psi_{\mathrm{fp32\ master}} + 4\Psi_{\mathrm{momentum}} + 4\Psi_{\mathrm{variance}} = 16\Psi\ \mathrm{bytes}.

所以 1.5B 参数的模型光 model states 就要约 24GB。参数本体只是其中一小部分,真正重的是 Adam 的 fp32 master weights、momentum、variance,以及梯度。

2. 已有解决方案与不足

已有路线主要有四类:

  1. Data parallelism:每个 rank 都有完整模型副本,训练 step 结束时 all-reduce gradients。优点是计算粒度大、通信模式简单、开发体验好;显存问题在于每张卡都保存同一份 parameters、gradients、optimizer states。
  2. Model parallelism:把每层参数或算子横向切到多张卡上。优点是单卡显存下降;代价是每层都需要通信,跨节点后 interconnect 带宽不足会快速降低效率。论文报告 40B 参数 Megatron-LM 跨两个 DGX-2 节点时,每张 V100 只有约 5 TFLOPs,低于硬件峰值 5%。
  3. Pipeline parallelism:按层纵向切分模型,并用 micro-batch 填充流水线。它能降低单阶段显存,但会引入 pipeline bubble、micro-batch 约束、activation 存储压力和模型改写成本。
  4. CPU offload / activation checkpointing / memory-efficient optimizer:这些方法分别降低某一类状态的显存,但 CPU 带宽、重计算、优化器语义变化或实现复杂度会成为限制。

作者看到的关键空白是:DP 的高效率和 MP 的低显存各自成立,但二者耦合了“哪里存状态”和“哪里做计算”。如果能把状态存储位置和计算执行位置解耦,训练时只在需要某个状态的时间窗口里通信和物化它,就能把 aggregate memory 用起来,同时保持大粒度 forward/backward。

3. 作者可能的思考路径

一个可能的推理路径如下。

第一步,作者从内存账本出发。单看 fp16 parameters 会低估训练显存,因为 Adam 的 fp32 master weights、momentum、variance 和 gradients 加起来远大于参数本体。DP 下这些状态在所有 rank 上复制,冗余倍数正好等于 data-parallel degree。

第二步,作者注意到 optimizer update 本身有分区性质。一个 data-parallel rank 更新某个参数分片时,只需要该分片对应的 reduced gradients 和 optimizer states。完整 gradients 在所有 rank 上常驻、完整 optimizer states 在所有 rank 上常驻,都只是传统 DP 实现的惯例。

第三步,作者把 all-reduce 拆成 reduce-scatter + all-gather。标准 DP 的 gradient all-reduce 可以看成先 reduce-scatter,再 all-gather。若每个 rank 只负责一个参数分片,reduce-scatter 后 rank 已经拿到自己需要更新的梯度;更新后再 all-gather 最新参数给所有 rank,通信量仍是 2Ψ2\Psi

第四步,作者进一步把 parameters 也分片。每层 forward/backward 只在该层需要参数时物化完整参数,用完后丢弃;这把长期常驻参数状态变成短生命周期通信缓存。额外代价是 forward 和 backward 都要 all-gather 参数,整体通信量从 2Ψ2\Psi 增到 3Ψ3\Psi,换来 model states 近似按 NdN_d 线性下降。

第五步,model states 降下来后,residual memory 成为新瓶颈。activations、large fused buffers、fragmentation 都会在超大模型上再次触发 OOM,所以 ZeRO-R 补上 activation partitioning、constant-size buffers 和 memory defragmentation。

这条路径的直觉很系统:先做精确内存分类,再找冗余,再利用训练时序,把“长期保存完整状态”改成“分片保存 + 按需重建”。

4. 核心假设或切入点

ZeRO 的核心假设有三个:

  1. 大模型训练显存的主要可优化对象是 model-state redundancy,尤其是 adaptive optimizer states。
  2. DP 的通信 collective 可以重新调度,使状态分片降低显存,同时保持通信 volume 接近标准 DP。
  3. 参数、梯度、activation checkpoint 都具有时间局部性。它们只在 forward/backward/update 的特定阶段需要完整物化,其他时间可以被分片保存或释放。

5. 方法 / 系统 / 理论框架

5.1 Memory accounting

论文用 Ψ\Psi 表示参数个数,用 KK 表示 optimizer states 的额外 memory multiplier。对于 mixed-precision Adam,K=12K=12,model states 总内存是:

MDP=4Ψ+KΨ=16Ψ. M_{\mathrm{DP}} = 4\Psi + K\Psi = 16\Psi.

这里 4Ψ4\Psi 对应 fp16 parameters 和 fp16 gradients,KΨK\Psi 对应 fp32 master weights、momentum、variance。

5.2 ZeRO-DP Stage 1: optimizer state partitioning PosP_{os}

每个 data-parallel rank 只保存 1/Nd1/N_d 的 optimizer states,并只更新自己负责的参数分片。更新完成后,通过 all-gather 得到完整更新参数。

这里通信的对象是 gradients 和更新后的 fp16 parameter shards,optimizer states 本身不在每个 step 做全量通信。rank rr 长期保存自己负责参数片对应的 fp32 master weights、momentum、variance;backward 后通过 gradient reduce-scatter 拿到这片参数的 averaged gradient shard,本地完成 Adam update,再把更新后的 parameter shard all-gather 给其他 ranks。

更具体地说,每个参数 shard 可以看成一个有状态的 optimizer instance。它拥有该 shard 的 master weights、mmvv,每步只接收对应的 averaged gradient shard,然后独立推进本地 Adam 状态。其他 ranks 只需要最终更新后的 fp16 parameter shard 来参与下一轮 forward/backward。

显存变为:

MPos=4Ψ+KΨNd. M_{P_{os}} = 4\Psi + \frac{K\Psi}{N_d}.

NdN_d 足够大、K=12K=12 时,极限近似为 4Ψ4\Psi,相对 16Ψ16\Psi 是约 4x model-state memory reduction。

直觉是:Adam 最重的状态被切分了,但 fp16 parameters 和 gradients 仍然完整常驻。

5.3 ZeRO-DP Stage 2: gradient partitioning Pos+gP_{os+g}

因为每个 rank 只更新自己的参数分片,它只需要该分片对应的 reduced gradients。于是 backward 中梯度产生后,系统执行 reduce-scatter,把不同参数分片的 reduced gradients 交给不同 rank,并释放完整梯度。

显存变为:

MPos+g=2Ψ+(K+2)ΨNd. M_{P_{os+g}} = 2\Psi + \frac{(K+2)\Psi}{N_d}.

对于 Adam,即:

MPos+g=2Ψ+14ΨNd. M_{P_{os+g}} = 2\Psi + \frac{14\Psi}{N_d}.

NdN_d 足够大时,极限近似为 2Ψ2\Psi,对应约 8x model-state memory reduction。

直觉是:此时长期常驻的完整 model state 只剩 fp16 parameters,gradient 和 optimizer states 都按 DP rank 分片。

5.4 ZeRO-DP Stage 3: parameter partitioning Pos+g+pP_{os+g+p}

每个 rank 也只长期保存 1/Nd1/N_d 的 parameters。某一层 forward/backward 需要完整参数时,负责该分片的 rank broadcast 或 all-gather 参数,所有 rank 在该层计算完成后丢弃非本地分片。

显存变为:

MPos+g+p=16ΨNd. M_{P_{os+g+p}} = \frac{16\Psi}{N_d}.

这就是 ZeRO 最强的 scaling 结论:model-state memory 随 data-parallel degree 线性下降。论文表格中,1T 参数模型的 model states 在 Nd=1024N_d=1024 时降到约 15.6GB/GPU,理论上能放进 32GB V100 的 model-state 预算。

5.5 Communication analysis

标准 DP 的 gradient all-reduce 可看作:

reduce-scatter(Ψ)+all-gather(Ψ), \mathrm{reduce\text{-}scatter}(\Psi) + \mathrm{all\text{-}gather}(\Psi),

所以每步通信 volume 约为:

CDP=2Ψ. C_{\mathrm{DP}} = 2\Psi.

PosP_{os} 中,通信调度可以复用这个分解:gradient reduce-scatter 需要 Ψ\Psi,更新后的 parameter shard all-gather 需要 Ψ\Psi,所以:

CPos=2Ψ. C_{P_{os}} = 2\Psi.

如果朴素地先做完整 gradient all-reduce,再额外 all-gather 更新后的参数,通信会变成 3Ψ3\Psi。论文的 2Ψ2\Psi 分析对应的是优化后的调度:省掉标准 all-reduce 的 gradient all-gather 半段,把这部分通信预算用于 all-gather updated parameters。

Pos+gP_{os+g} 中,梯度 reduce-scatter 需要 Ψ\Psi,参数更新后的 all-gather 需要 Ψ\Psi,所以:

CPos+g=2Ψ. C_{P_{os+g}} = 2\Psi.

Pos+g+pP_{os+g+p} 中,参数在 forward 和 backward 各需要按层 all-gather 一次,加上 gradient reduce-scatter:

CPos+g+p=Ψforward param+Ψbackward param+Ψgrad reduce-scatter=3Ψ. C_{P_{os+g+p}} = \Psi_{\mathrm{forward\ param}} + \Psi_{\mathrm{backward\ param}} + \Psi_{\mathrm{grad\ reduce\text{-}scatter}} = 3\Psi.

因此完整 stage 3 的通信 volume 是标准 DP 的 1.5x。论文强调这是用可接受的通信增加换取按 NdN_d 线性下降的 model-state memory。

5.6 ZeRO-R: residual memory optimizations

当 model states 被 ZeRO-DP 压低后,剩余显存瓶颈来自 activations、temporary buffers 和 fragmentation。

ZeRO-R 包含三部分:

  1. Partitioned activation checkpointing PaP_a:在 model parallelism 下,activation checkpoints 常被多个 MP rank 复制。ZeRO 把 checkpointed activations 按 MP degree 分片存储,在 backward recomputation 前 all-gather 重建。对于 100B 模型、batch size 32、sequence length 1024、MP=16,论文估计 activation checkpoints 可从约 33GB/GPU 降到约 2GB/GPU。
  2. CPU activation checkpoint offload Pa+cpuP_{a+cpu}:在极大模型或显存极紧时,把 partitioned activation checkpoints 放到 CPU,换取额外 CPU-GPU 数据移动。
  3. Constant-size buffers CBC_B 与 memory defragmentation MDM_D:避免 large fused buffer 随模型大小线性增长;把 activation checkpoints 和 gradients 移入预分配连续 buffer,降低碎片导致的 OOM 和 allocator overhead。

6. 结论链条

论文的结论链条可以压缩为:

  1. 大模型训练显存主要被 model states 占据,mixed-precision Adam 需要约 16Ψ16\Psi bytes。
  2. 标准 DP 在每个 rank 上复制完整 model states,memory redundancy 随 DP degree 增长。
  3. Optimizer states、gradients、parameters 可以累积分片,得到 PosP_{os}Pos+gP_{os+g}Pos+g+pP_{os+g+p} 三个阶段。
  4. Pos+gP_{os+g} 可保持与 DP 相同的 2Ψ2\Psi 通信 volume,同时把 model-state memory 极限降到约 2Ψ2\Psi
  5. Pos+g+pP_{os+g+p} 把 model-state memory 降到 16Ψ/Nd16\Psi/N_d,通信 volume 增至 3Ψ3\Psi
  6. ZeRO-R 处理 activations、buffers、fragmentation,让 model states 降低后产生的新瓶颈继续可控。
  7. 实测 ZeRO-100B 在 400 张 V100 上能高效训练 100B+ GPT-like 模型,并将最大可运行模型推到 170B。

关键实验/定理

结果 1:内存公式和最大模型规模分析

  • 设置:mixed-precision Adam,K=12K=12,比较 DP degree NdN_d 从 1 到 1024 时,7.5B、128B、1T 模型在三阶段 ZeRO-DP 下的 per-device model-state memory。
  • 指标:每张 GPU 的 model-state memory GB。
  • 结果:7.5B 模型在 Nd=64N_d=64 时,标准 DP 需要 120GB;PosP_{os} 约 31.4GB;Pos+gP_{os+g} 约 16.6GB;Pos+g+pP_{os+g+p} 约 1.88GB。1T 模型在 Nd=1024N_d=1024Pos+g+pP_{os+g+p} 下约 15.6GB/GPU。
  • 解读:ZeRO 的核心分析聚焦把 DP 的 aggregate memory 转化为大模型可用的 model-state memory,节省显存只是这个系统目标的直接结果。完整 stage 3 的理论意义尤其大。

结果 2:通信 volume 分析

  • 设置:比较标准 DP、Pos+gP_{os+g}Pos+g+pP_{os+g+p} 每步通信量。
  • 指标:相对于参数规模 Ψ\Psi 的通信 volume。
  • 结果:标准 DP 为 2Ψ2\PsiPos+gP_{os+g} 仍为 2Ψ2\PsiPos+g+pP_{os+g+p}3Ψ3\Psi,即 1.5x 标准 DP。
  • 解读:ZeRO 的系统价值在于 memory saving 和 communication volume 没有线性绑定。Stage 1/2 基本保留 DP 通信量;Stage 3 用 50% 通信增加换取按 DP degree 分摊 model states。

结果 3:ZeRO-100B 在 400 V100 上训练 100B+ 模型

  • 设置:实现 ZeRO-100B,即 Pos+gP_{os+g} 加 ZeRO-R;硬件为 400 张 V100、25 个 DGX-2 节点、800Gbps internode bandwidth;模型为 GPT-2-like Transformer;baseline 为 Megatron-LM 2019 年开源版本。
  • 指标:最大可运行模型规模、per-GPU throughput、aggregate throughput。
  • 结果:ZeRO-100B 可高效运行最高 170B 参数模型;100B 模型达到超过 38 TFLOPs/GPU,总吞吐超过 15 PFLOPs;相比 baseline 大模型设置最高约 10x speedup。
  • 解读:实测证明前两阶段分片加 residual memory optimizations 已经足够把当时可训练规模从 10B 级推进到 100B 级。完整 stage 3 的 trillion claim 仍主要是分析结果。

结果 4:super-linear scalability

  • 设置:60B 参数模型,GPU 数从 64 增至 400,MP degree 固定为 16。
  • 指标:吞吐 scaling 和 per-GPU training throughput。
  • 结果:论文观察到 64 到 400 GPU 区间的 super-linear speedup。
  • 解读:原因是 DP degree 提高后,ZeRO-DP 降低每卡 model-state memory,使每卡能放更大 batch。更大 batch 提高 arithmetic intensity,吞吐随 GPU 数增加超过线性预期。这个结论依赖 batch size 仍处在不会明显损害收敛的区间。

结果 5:无需 MP 训练 13B 模型

  • 设置:128 张 V100,只用 ZeRO-powered DP,不用 model parallelism 或 pipeline parallelism。
  • 指标:最大可训练模型规模与 throughput。
  • 结果:ZeRO-100B 可训练最高 13B 参数模型;传统 DP baseline 在约 1.4B 参数附近 OOM。
  • 解读:这是论文的 usability 结果。许多研究者无需写 MP/PP 版本模型,也能探索超过当时 T5 11B 级别的模型规模。

结果 6:Turing-NLG 17B

  • 设置:Turing-NLG 17B 使用 ZeRO-100B 端到端训练。
  • 指标:WebText-103 perplexity、training throughput。
  • 结果:论文称截至 2020-05-12,Turing-NLG 是 17B+ 参数 language model,在 WebText-103 上 perplexity 10.21,训练吞吐 41.4 TFLOPs/GPU。
  • 解读:这是 ZeRO 从系统原型进入真实大模型训练的示例。它证明 ZeRO 不只提高 synthetic GPT-like benchmark 的可运行规模,也支撑了当时公开发布的 Microsoft 大模型。

证据链强度评估

强证据

  • Memory accounting 很强。mixed-precision Adam 的 16Ψ16\Psi 分解清楚,三阶段分片公式可直接复算。
  • Pos+gP_{os+g} 的通信 volume 分析强。把 all-reduce 拆成 reduce-scatter + all-gather 后,通信量与标准 DP 相同,这个推导简洁且工程上可实现。
  • ZeRO-100B 实测覆盖 400 张 V100、最高 170B 模型,实验规模在 2020 年很强。

中等强度证据

  • Super-linear scalability 的解释合理,但依赖 batch size、模型规模、硬件拓扑和收敛区间。换到不同模型、数据和 optimizer schedule 后需要重新验证。
  • ZeRO-R 的 activation partitioning / CPU offload trade-off 具有清楚机制,但 CPU offload 是否提升吞吐取决于 CPU-GPU 带宽和可增加 batch size 的收益。
  • 与 Megatron-LM baseline 的对比有历史价值,但 baseline 是 2019 开源版本,当前结论更适合作为系统演化坐标。

需要谨慎的推论

  • 1T 参数训练在本文中主要是 memory feasibility,不等于端到端训练时间已经可接受。论文自己指出 1T 模型在 1024 GPU 级别可能需要 140 天到一年以上,实际需要 exa-flop 级系统。
  • Stage 3 在 v3 论文中的重点是分析,实测实现范围是 Pos+gP_{os+g} 加 ZeRO-R。后续 DeepSpeed ZeRO-3 的工程成熟度属于后续系统演进。
  • ZeRO 处理 dense Transformer 的 model-state redundancy。MoE expert routing、expert parallel all-to-all、long-context KV cache、RL rollout/trainer mismatch 等后续瓶颈需要其他系统论文补齐。

OpenReview / 审稿意见吸收

  • Venue status: 当前档案未记录公开 peer-review 状态。
  • Public reviews: 当前档案未记录可可靠匹配的 OpenReview / ARR / 会议 reviewer comments。
  • Ratings / confidence: 无公开评分可用于校准。
  • Reviewer consensus: 暂无。
  • Main criticisms: 暂无公开 reviewer 质疑可引用;可信度主要由论文、技术报告、项目证据和本地一致性检查决定。
  • Author response: 暂无公开 rebuttal 记录。
  • 对本文可信度的影响: 按未完成公开审稿吸收处理,结论需要依赖实验设置、baseline 强度、复现证据和跨论文一致性校准。

本地讨论补充

1. 讨论收敛点

  • 当前初版阅读将 ZeRO 定位为本地 archive 中“训练内存系统”的基础节点。它解释了为什么后续大模型训练框架会把 optimizer/gradient/parameter sharding 当作默认组件。
  • ZeRO 和 FlashAttention 的互补关系很清楚:ZeRO 降低跨 GPU 训练状态冗余,FlashAttention 降低单 kernel 内 attention IO。二者都把“显存容量/带宽”上升为模型 scaling 的一等变量。
  • ZeRO 和 HybridFlow/VERL/slime 的关系是层级关系:后者编排 RLHF/RLVR 多模型 dataflow、rollout 和 policy update;ZeRO/FSDP/Megatron 这类底层机制决定 actor/ref/reward/trainer 能用多大模型和多大 batch。

2. 修正后的理解

  • ZeRO 的名称虽然叫 Optimizer,但论文贡献覆盖 optimizer states、gradients、parameters、activation checkpoints、temporary buffers 和 fragmentation。Optimizer state partitioning 是第一阶段,完整系统远大于一个 optimizer wrapper。
  • Stage 1/2 的关键点是“用分片状态替换复制状态,同时让通信量保持在 DP 的 2Ψ2\Psi 量级”。Stage 3 的关键点是“参数本体也不再长期完整常驻”,用按层 all-gather 换取 16Ψ/Nd16\Psi/N_d model-state memory。
  • ZeRO 的 trillion claim 应写成“model-state memory feasibility”,并和“训练时间/算力 feasibility”分开。论文明确指出 1T 模型的端到端训练仍需要远超当时集群的算力。

3. 为什么优化器状态可以分片

在标准 data parallelism 中,每个 rank 处理不同 mini-batch shard,然后通过 gradient all-reduce 得到同一份全局平均梯度。由于每个 rank 拿到相同参数、相同平均梯度、相同 optimizer hyperparameters,它们会执行完全相同的 Adam 更新。因此每张 GPU 上的 Adam states 是重复副本:fp32 master weights、momentum mm、variance vv 都按参数坐标一一对应。

以 Adam 的单个参数坐标 ii 为例:

miβ1mi+(1β1)gi, m_i \leftarrow \beta_1 m_i + (1-\beta_1) g_i,
viβ2vi+(1β2)gi2, v_i \leftarrow \beta_2 v_i + (1-\beta_2) g_i^2,
θiθiηm^iv^i+ϵ. \theta_i \leftarrow \theta_i - \eta \frac{\hat m_i}{\sqrt{\hat v_i}+\epsilon}.

这个更新只依赖 θi\theta_igig_imim_iviv_i 和超参,不依赖其他参数坐标的 optimizer state。于是可以把参数坐标集合切成 NdN_d 份:rank 0 只保存第 0 份参数的 optimizer states,rank 1 只保存第 1 份,以此类推。每个 rank 只负责自己那一片参数的 optimizer update。

实现顺序可以按 ZeRO stage 理解:

  1. PosP_{os}:optimizer states 分片。每个 rank 保存自己负责参数片的 fp32 master weights、momentum、variance。backward 后通过 reduce-scatter 拿到自己负责参数片的 averaged gradient shard,然后本地更新这一片参数。更新完成后 all-gather 最新参数,让所有 rank 重新拥有完整 fp16 parameters,供下一轮 forward/backward 使用。optimizer states 长期留在 owner rank,本身不做 step 级全量通信。
  2. Pos+gP_{os+g}:gradient 也分片。backward 过程中每个 gradient bucket 产生后,用 reduce-scatter 把对应参数片的 reduced gradients 分给负责它的 rank。这样 rank 只保留自己要更新的 gradient shard,更新后再 all-gather 参数。通信量仍约为标准 DP 的 2Ψ2\Psi:一次 reduce-scatter 加一次 all-gather。
  3. Pos+g+pP_{os+g+p}:parameters 也长期分片。rank 长期只保存自己那一片参数;forward/backward 到某一层时,再按需 all-gather 该层完整参数,用完释放非本地分片。这样 model states 变成 16Ψ/Nd16\Psi/N_d,通信量增加到约 3Ψ3\Psi

可以把它形象化为:标准 DP 是每个工人都保存完整账本,并且每个人都重复计算整本账的更新;ZeRO 把账本按页分给不同工人,每个人只维护自己负责的页。训练计算需要整本账时临时复印相关页,更新账本时各自改自己那一页,然后同步最新版。

4. 为什么 all-reduce 可以拆成 reduce-scatter 和 all-gather

All-reduce 的语义是:所有 rank 输入同样长度的 tensor,系统先对同一位置做 sum / average,然后让每个 rank 都拿到完整 reduced tensor。

假设有 NdN_d 个 rank,每个 rank 的 gradient vector 长度为 Ψ\Psi。把 gradient vector 切成 NdN_d 个连续 shard:

g=[g(0),g(1),,g(Nd1)]. g = [g^{(0)}, g^{(1)}, \ldots, g^{(N_d-1)}].

All-reduce 的最终结果是每个 rank 都得到:

gˉ=[gˉ(0),gˉ(1),,gˉ(Nd1)], \bar g = [ \bar g^{(0)}, \bar g^{(1)}, \ldots, \bar g^{(N_d-1)} ],

其中:

gˉ(k)=1Ndr=0Nd1gr(k). \bar g^{(k)} = \frac{1}{N_d}\sum_{r=0}^{N_d-1} g_r^{(k)}.

这个语义可以自然拆成两个阶段:

  1. Reduce-scatter:对每个 shard 做跨 rank reduce,然后把第 kk 个 reduced shard 放到负责它的 rank kk 上。结束后,每个 rank 只拥有一段全局平均后的 gradient shard。
  2. All-gather:把各 rank 手里的 reduced shards 再收集到所有 rank。结束后,每个 rank 都拥有完整 gˉ\bar g

所以标准 DP 的 all-reduce 可以理解为:

all-reduce(g)=all-gather(reduce-scatter(g)). \mathrm{all\text{-}reduce}(g) = \mathrm{all\text{-}gather}( \mathrm{reduce\text{-}scatter}(g) ).

这也是很多高性能 ring all-reduce 的实际算法结构:第一圈做 reduce-scatter,第二圈做 all-gather。对大 tensor 来说,每个 rank 在 reduce-scatter 阶段移动约 Nd1NdΨ\frac{N_d-1}{N_d}\Psi 个元素,在 all-gather 阶段再移动约 Nd1NdΨ\frac{N_d-1}{N_d}\Psi 个元素,总量约为 2Ψ2\Psi

ZeRO 利用的是 reduce-scatter 结束后的中间态。标准 DP 需要完整 reduced gradients,因为每个 rank 都要更新完整参数和完整 optimizer states。ZeRO-2 中,每个 rank 只负责更新一个 parameter shard,因此它只需要对应的 gradient shard。这样训练可以在 reduce-scatter 后直接进入 optimizer update,避免每个 rank 长期保存完整 gradient。

ZeRO-2 仍然需要让下一步 forward 使用完整 fp16 parameters,所以它把 all-gather 放到 optimizer update 之后,用来同步更新后的 parameter shards:

local gradsreduce-scatterreduced grad shardsoptimizer updateupdated param shardsall-gatherfull fp16 params. \text{local grads} \xrightarrow{\mathrm{reduce\text{-}scatter}} \text{reduced grad shards} \xrightarrow{\mathrm{optimizer\ update}} \text{updated param shards} \xrightarrow{\mathrm{all\text{-}gather}} \text{full fp16 params}.

通信量仍接近标准 DP 的 2Ψ2\Psi,但显存占用明显不同:标准 DP 的 all-reduce 之后每个 rank 都有完整 gradients;ZeRO-2 的 reduce-scatter 之后每个 rank 只有自己负责的 gradient shard。

5. 为什么把这些 collective 作为基础操作

ZeRO 把 all-reduce、reduce-scatter、all-gather 当作基础操作,原因是它们正好覆盖分布式训练中三种最常见的 tensor 状态转换。

  1. Replicated local tensor 到 replicated reduced tensor:每个 rank 都有一份 local gradient,需要所有 rank 得到同一份全局平均 gradient。这个转换对应 all-reduce,是标准 DP 的核心操作。
  2. Replicated local tensor 到 sharded reduced tensor:每个 rank 都有一份 local gradient,但每个 rank 只需要自己负责参数片的全局平均 gradient。这个转换对应 reduce-scatter,是 ZeRO-2/3 更新 gradient shard 的核心操作。
  3. Sharded tensor 到 replicated full tensor:每个 rank 只持有一个 parameter shard,但 forward/backward 需要某层完整参数。这个转换对应 all-gather,是 ZeRO-1/2 step 末尾同步参数、ZeRO-3 按层临时重建参数的核心操作。

换句话说,ZeRO 的状态管理只有三个基本问题:

  • 需要把多份局部结果合成一份全局结果:reduce。
  • 需要让每个 rank 只保留自己负责的那一片:scatter。
  • 需要把分片状态重新拼成完整视图:gather。

All-reduce、reduce-scatter、all-gather 刚好是这些动作在多 GPU 上的标准化组合。用这些 collective 作为基础操作有几个工程好处:

  1. 语义和 sharding 对齐。参数、梯度、optimizer states 都按同一套 shard 边界切分,reduce-scatter 和 all-gather 可以直接作用在 shard 上。
  2. 通信量容易分析。标准 DP 是 2Ψ2\Psi,ZeRO-2 仍是一次 reduce-scatter 加一次 all-gather,ZeRO-3 增加 forward/backward parameter all-gather 后到约 3Ψ3\Psi
  3. 实现有成熟优化。NCCL、NVLink、InfiniBand、ring/tree collective 都长期优化这些 primitive,大 tensor 训练通常受 bandwidth 主导,标准 collective 更容易接近硬件带宽。
  4. 可以和计算重叠。Backward 按 layer / bucket 产生梯度,reduce-scatter 可以在梯度 bucket 产生后立刻启动;ZeRO-3 的 parameter all-gather 也可以按 layer 预取,减少等待。
  5. 它们可组合。标准 all-reduce 可以看作 reduce-scatter + all-gather;ZeRO 保留 reduce-scatter 的中间态做 optimizer update,再把 all-gather 移到参数同步或参数重建的位置。

所以这些操作被选为基础操作,主要因为它们是分布式张量在 full / shard / reduced 三种状态之间移动的最小通用接口。ZeRO 的创新点在于改变这些接口出现的位置:标准 DP 先 all-reduce 出完整 gradients,再完整更新;ZeRO 在 reduce-scatter 后停留在 shard 状态,局部更新 optimizer states,再按需要 all-gather 参数。

6. w/o ZeRO 时实际 LLM 如何同步各层参数

这里要区分“参数同步”和“梯度同步”。在普通 data parallel / DDP 训练里,每个 rank 从 step 开始就持有完整且相同的参数副本。Forward 经过 embedding、每个 Transformer block、LM head 时,参数已经在本地,不需要每层 all-gather 参数。真正发生跨 rank 通信的位置通常在 backward:各层梯度产生后,通过 gradient all-reduce 同步。

一个典型 LLM DDP step 可以这样理解:

  1. 初始化或加载 checkpoint:rank 0 或 checkpoint loader 将同一份模型权重广播 / 加载到所有 data-parallel ranks。此后每个 rank 都有完整 embedding、所有 Transformer layers、LM head 参数。
  2. Forward:rank rr 处理自己的 micro-batch shard。Embedding lookup、Attention、MLP、RMSNorm/LayerNorm、LM head 都使用本 rank 本地完整参数。纯 DP 情况下,forward 不需要跨 DP ranks 同步参数。
  3. Loss:每个 rank 得到自己的 local loss,通常对应不同样本 shard。
  4. Backward 从 LM head 向前传播:LM head gradients 先产生,接着是最后一个 Transformer block 的 MLP、attention、norm gradients,再按层反向直到 embedding。
  5. DDP autograd hooks:每个 parameter gradient ready 后,DDP 把它标记到某个 gradient bucket。bucket 常按内存大小组织,可能包含一个层的一部分参数,也可能跨多个相邻层;实际通信单位通常是 bucket,层边界和 bucket 边界经常不同。
  6. Bucket all-reduce:当某个 bucket 里的 gradients 都 ready 后,DDP 立即对这个 bucket 启动 all-reduce。这个 all-reduce 在高性能实现里常由 reduce-scatter + all-gather 组成:先把 bucket 切块并归约到不同 ranks,再把 reduced chunks 收集回所有 ranks。完成后,每个 rank 都拥有相同的 averaged gradient bucket。
  7. Backward/communication overlap:当较后层的 bucket 在 all-reduce 时,backward 还可以继续计算更前层 gradients。这样通信和反向计算重叠,降低等待。
  8. Optimizer step:所有 buckets 的 all-reduce 完成后,每个 rank 都有完整 averaged gradients。由于每个 rank 的参数、optimizer states、averaged gradients 和超参相同,AdamW/Adam update 结果也相同。
  9. Step 结束:每个 rank 的完整参数副本继续保持一致。下一步 forward 直接使用本地参数。

因此 w/o ZeRO 的“参数同步”主要是间接发生的:

local layer gradsbucket all-reducereplicated averaged gradssame optimizer updatereplicated updated params. \text{local layer grads} \xrightarrow{\mathrm{bucket\ all\text{-}reduce}} \text{replicated averaged grads} \xrightarrow{\mathrm{same\ optimizer\ update}} \text{replicated updated params}.

在这个流程中,all-reduce 同步的是 gradients;parameters 通过相同 optimizer update 保持同步。DDP 可能在初始化、异常恢复或参数校验时广播参数,但每个普通 training step 里,纯 data parallel 的主要通信路径是 backward gradient all-reduce。

如果 LLM 还使用 tensor parallelism,layer 内部会有额外 activation / partial-output collectives,例如 attention/MLP linear 的 all-reduce 或 all-gather;这些属于 tensor-parallel layer computation,和 w/o ZeRO 的 data-parallel 参数同步是两类问题。ZeRO 论文中讨论的 DP model-state redundancy,主要针对 data-parallel ranks 之间完整复制参数、梯度和 optimizer states 的问题。

7. w/ ZeRO 时实际 LLM 如何同步各层参数

w/ ZeRO 后,同样要区分 data-parallel 维度和 tensor-parallel 维度。ZeRO 处理的是 data-parallel ranks 之间的 model states 放置:optimizer states、gradients、parameters 是否长期复制。它不改变 Transformer layer 的数学形式;每个 DP rank 仍然处理自己的 batch shard,attention、MLP、norm、LM head 的局部计算顺序保持一致。变化发生在状态何时完整存在、何时分片保存、何时用 collective 临时重建。

ZeRO-1 的通信路径如下:

  1. Step 开始:每个 rank 拥有完整 fp16 parameters;optimizer states 按参数坐标分片。
  2. Forward:和普通 DDP 接近,每层参数已经在本地,embedding、Transformer blocks、LM head 不需要按层 gather 参数。
  3. Backward:每个 rank 得到 local gradients。优化后的通信调度对 gradient bucket 做 reduce-scatter,把 averaged gradient shard 交给负责对应参数片的 rank。
  4. Optimizer update:rank rr 只用自己负责参数片的 reduced gradients 更新本地 fp32 master weights、momentum、variance。
  5. Parameter sync:各 rank 把更新后的 parameter shards all-gather 给所有 rank,使下一步 forward 重新拥有完整 fp16 parameters。

这样 ZeRO-1 的 step 级通信仍是 2Ψ2\Psi:一次 gradient reduce-scatter 加一次 updated-parameter all-gather。朴素实现若先做完整 gradient all-reduce,再同步更新后的参数,会额外多出一次 parameter all-gather。

这也是 ZeRO-1 和普通 DDP 最关键的状态所有权差异:普通 DDP 中每个 rank 都有一套完整 optimizer instances,并对所有参数坐标重复执行同一更新;ZeRO-1 中每个 optimizer instance 只存在于 owner rank,对应 shard 的梯度被路由过来,本地状态独立更新,随后广播的是更新后的参数值。

ZeRO-2 的关键变化是 gradient residency。它把普通 DDP 的 gradient all-reduce 替换成 reduce-scatter + parameter all-gather:

  1. Forward:每个 rank 仍有完整 fp16 parameters,所以每层 forward 不需要 data-parallel parameter all-gather。
  2. Backward:梯度按 bucket 产生。bucket ready 后立即 reduce-scatter,rank rr 只收到自己负责参数片的 averaged gradient shard。
  3. Optimizer update:rank rr 用本地 gradient shard 更新本地 optimizer-state shard 和 master-weight shard。
  4. Parameter sync:更新后的 parameter shards all-gather,所有 rank 恢复完整 fp16 parameters。

因此 ZeRO-2 和 w/o ZeRO 的每层计算看起来很像,主要差异在 backward bucket 完成后的状态形态:

w/o ZeRO: local grad bucketall-reducereplicated averaged grad bucketlocal Adamreplicated updated params. \text{w/o ZeRO: local grad bucket} \xrightarrow{\mathrm{all\text{-}reduce}} \text{replicated averaged grad bucket} \xrightarrow{\mathrm{local\ Adam}} \text{replicated updated params}.
w/ ZeRO-2: local grad bucketreduce-scatteraveraged grad shardsharded Adamupdated param shardall-gatherreplicated updated params. \text{w/ ZeRO-2: local grad bucket} \xrightarrow{\mathrm{reduce\text{-}scatter}} \text{averaged grad shard} \xrightarrow{\mathrm{sharded\ Adam}} \text{updated param shard} \xrightarrow{\mathrm{all\text{-}gather}} \text{replicated updated params}.

ZeRO-3 进一步把 parameters 也变成长期分片状态。它的同步点进入 layer 级别:

  1. Step 开始:每个 rank 只长期保存自己负责的 parameter shard、optimizer-state shard;gradient shard 在 backward/update 阶段产生和释放。
  2. Layer forward 前:对当前 layer 的 parameter shards 做 all-gather,让每个 rank 临时拥有该 layer 完整参数。
  3. Layer forward 后:释放非本地 parameter shards,降低峰值和常驻显存。
  4. Layer backward 前:反向经过该 layer 时,再次 all-gather 该 layer 参数,因为 gradient 计算也需要对应权重。
  5. Gradient reduce-scatter:该 layer 的 gradients 产生后,直接 reduce-scatter 到负责对应参数片的 rank。
  6. Optimizer update:rank rr 更新自己的 parameter shard 和 optimizer-state shard。step 结束后仍保持 shard-only 状态,下一步 forward 再按 layer gather。

所以 w/o ZeRO 和 w/ ZeRO 的核心对照是状态生命周期:

场景 参数常驻状态 梯度同步方式 Optimizer states 参数同步位置
w/o ZeRO / DDP 每个 rank 完整复制 bucket all-reduce,完成后每个 rank 有完整 averaged gradients 每个 rank 完整复制 通过相同 optimizer update 间接保持一致
ZeRO-1 每个 rank 完整复制 通常仍接近 all-reduce / shard update 按参数坐标分片 step 末尾 all-gather 更新后的 parameter shards
ZeRO-2 每个 rank 完整复制 bucket reduce-scatter,只保留本 rank gradient shard 按参数坐标分片 step 末尾 all-gather 更新后的 parameter shards
ZeRO-3 长期只保存 shard,按 layer 临时重建 layer/bucket reduce-scatter 按参数坐标分片 forward/backward 前按 layer all-gather,step 末尾保持 shard

如果同时启用 tensor parallelism,ZeRO 的 data-parallel state sharding 会和 tensor-parallel layer collectives 叠加。例如一个 MLP linear 可能先在 TP 组内做 row/column-parallel 的 all-reduce 或 all-gather,同时 ZeRO-3 在 DP 组内为该 layer gather parameter shards。两组通信的 process group、切分语义和生命周期不同:TP 通信服务 layer 内部计算,ZeRO 通信服务 DP model-state residency。

8. Collective 通信实际按什么维度切

集合通信本身通常不关心“feature 维”“sequence 维”这类模型语义。NCCL / MPI 看到的是一段 tensor buffer 或一组连续 chunks。具体切成哪种模型维度,取决于上层并行策略如何组织这段 buffer。

可以分三层理解:

  1. 通信库层:ring all-reduce、reduce-scatter、all-gather 会把输入 buffer 切成若干连续 chunks。这个切分主要服务带宽利用、pipeline 和 network topology,语义上接近“线性内存切块”。
  2. Data parallel / ZeRO 层:DDP 的 all-reduce 通常作用在 flattened gradient bucket 上,bucket 可能包含多个参数 tensor 的连续片段。ZeRO 的 reduce-scatter / all-gather 通常作用在 parameter / gradient / optimizer-state shards 上,也更接近 flattened parameter coordinates。这里的 shard 边界由参数分片和 bucketization 决定,经常跨 layer 或跨 tensor。
  3. Tensor parallel 层:这里才经常出现 feature 维、head 维、hidden 维切分。例如 column-parallel linear 按 output feature 切权重,row-parallel linear 按 input feature 切权重,attention 可以按 heads 切。对应的 all-gather / all-reduce 会带有明确模型语义。

所以如果讨论 ZeRO / DDP 的 data-parallel communication,切分对象主要是 parameter / gradient buffer,不一定对应 feature 维。一个 Transformer layer 的权重矩阵 WRdout×dinW\in\mathbb{R}^{d_{\mathrm{out}}\times d_{\mathrm{in}}} 进入 DDP bucket 后,通信库可能只看到 flatten 后的一段连续内存。ZeRO-3 gather 某层参数时,也是在收集该层 parameter shards;这些 shards 可以映射回矩阵的一段 row、column 或 flattened slice,具体取决于实现的 flattening 和 partition policy。

如果讨论 Megatron-style tensor parallel,feature 维切分就很常见:

  • Attention QKV / MLP up-projection 的 column-parallel:每个 rank 负责一部分 output channels / heads,后续可能需要 all-gather 或保持分片进入下一层。
  • MLP down-projection / output projection 的 row-parallel:每个 rank 负责一部分 input features,局部 matmul 后需要 all-reduce 合并 partial outputs。
  • Sequence / context parallel:切分对象会变成 sequence tokens 或 context blocks。
  • Expert parallel:切分对象会变成 experts、tokens-to-experts 或 expert load。

因此,“all-gather 等 collective 几乎都通过切 feature 维实现”只适用于一部分 tensor-parallel layer computation。ZeRO 论文关注的 DP/optimizer-state sharding,主要是按参数坐标和 flattened bucket 切;通信 primitive 相同,逻辑维度不同。

9. 一步训练里的 forward / backward / optimizer flow

以 mixed-precision Adam + data parallel training 为例,一步训练可以拆成四个状态流:fp16 parameters 用于 forward/backward,activations 用于 backward,gradients 用于 optimizer update,fp32 optimizer states 用于稳定更新。

标准 DP 的流程如下:

  1. Step 开始:每个 GPU 都有完整 fp16 parameters、完整 fp32 master weights、完整 Adam momentum mm、完整 Adam variance vv。不同 GPU 拿到不同 mini-batch shard。
  2. Forward:每个 GPU 用完整 fp16 parameters 处理自己的 batch shard,并保存 backward 需要的 activations 或 activation checkpoints。
  3. Backward:每个 GPU 根据本地 loss 反传,得到本地 gradients。此时 gradients 语义上仍是 local batch gradients。
  4. Gradient sync:所有 GPU 对 gradients 做 all-reduce,得到相同的 global averaged gradients。同步结束后,每个 GPU 都有完整 reduced gradients。
  5. Optimizer update:每个 GPU 使用完整 reduced gradients 更新完整 mmvv、fp32 master weights,再把更新后的 fp32 master weights cast / copy 回 fp16 parameters。
  6. Step 结束:每个 GPU 再次拥有完全相同的完整模型和完整 optimizer states,进入下一步。

这个流程的显存问题在第 1、4、5、6 步:每个 GPU 都长期保存同一份 optimizer states 和 parameters,并在 update 阶段保存完整 gradients。对于 Adam mixed precision,单个参数对应 fp16 parameter、fp16 gradient、fp32 master weight、fp32 momentum、fp32 variance,总计约 16 bytes。

ZeRO 的改写是逐步减少这些“完整常驻状态”。

Stage 1: PosP_{os} 只切 optimizer states。

  1. Step 开始:每个 GPU 仍有完整 fp16 parameters;每个 GPU 只保存自己那一片 fp32 master weights、mmvv
  2. Forward / backward:计算方式接近标准 DP,因为每个 GPU 有完整 fp16 parameters。
  3. Gradient sync:优化路径是 reduce-scatter,让每个 rank 拿到自己负责参数片的 averaged gradient shard;梯度 memory 仍按 Stage 1 公式保留完整梯度项,因为它没有像 Stage 2 那样把 gradient residency 系统化改成分片。
  4. Optimizer update:rank rr 只更新自己负责的 parameter shard 对应的 fp32 master weights、mmvv
  5. Parameter sync:各 rank 把自己更新后的 parameter shard all-gather 给其他 rank,使所有 GPU 恢复完整 fp16 parameters。

Stage 2: Pos+gP_{os+g} 继续切 gradients。

  1. Step 开始:每个 GPU 有完整 fp16 parameters;optimizer states 仍按参数分片。
  2. Forward:与标准 DP 接近。
  3. Backward:梯度按 bucket 产生。某个 bucket 的梯度一出来,就做 reduce-scatter,把 reduced gradient shard 直接发给负责该参数片的 rank。
  4. Gradient residency:每个 rank 只保留自己要更新的 gradient shard,其他 gradient 不长期保存。
  5. Optimizer update:每个 rank 用本地 gradient shard 更新本地 optimizer-state shard 和本地 fp32 master-weight shard。
  6. Parameter sync:更新后的 parameter shards all-gather,所有 rank 获得下一步 forward 所需的完整 fp16 parameters。

Stage 2 的关键变化是把标准 DP 的 all-reduce 拆成 reduce-scatter + all-gather。reduce-scatter 负责同步并分发 gradients,all-gather 负责同步更新后的 parameters。总通信量仍约为 2Ψ2\Psi

Stage 3: Pos+g+pP_{os+g+p} 继续切 parameters。

  1. Step 开始:每个 GPU 只长期保存自己的 parameter shard、gradient shard 和 optimizer-state shard。
  2. Layer forward 前:系统 all-gather 当前 layer 需要的 parameter shards,让所有 rank 临时拥有该 layer 的完整参数。
  3. Layer forward:所有 rank 用临时完整参数计算自己的 batch shard,并保存必要 activation checkpoint。
  4. Layer forward 后:释放非本地 parameter shards,只保留本 rank 长期负责的 shard。
  5. Backward:按反向 layer 顺序再次 all-gather 当前 layer 参数,计算 local gradients;梯度产生后 reduce-scatter 到负责对应参数片的 rank。
  6. Optimizer update:每个 rank 使用自己的 gradient shard 更新自己的 fp32 master weights、mmvv 和 parameter shard。
  7. Step 结束:parameters、gradients、optimizer states 都只以 shard 形式长期存在。下一步 forward 再按 layer 临时 all-gather。

可以用一句话概括三类状态的生命周期:

  • Parameters:标准 DP 中全程完整常驻;ZeRO-1/2 仍完整常驻;ZeRO-3 只在当前 layer forward/backward 时临时完整。
  • Gradients:标准 DP 中 all-reduce 后完整存在于每个 rank;ZeRO-2/3 用 reduce-scatter 让每个 rank 只保留自己负责的 gradient shard。
  • Optimizer states:标准 DP 中每个 rank 完整保存;ZeRO-1/2/3 都按参数 shard 长期保存。

因此 ZeRO 的本质是改变状态生命周期:计算需要完整视图时临时构造,更新需要局部状态时只保留局部分片。

10. ZeRO 改写后的流程总览

如果只看 ZeRO 改写后的训练 step,可以按时间线理解:

  1. 持久状态初始化:把参数坐标划分给不同 data-parallel ranks。ZeRO-1/2 中,每个 rank 长期保存完整 fp16 parameters,但 optimizer states 只保存本 rank 对应 shard;ZeRO-3 中,fp16 parameters、gradients、optimizer states 都只长期保存 shard。
  2. Forward 前准备参数:ZeRO-1/2 已经有完整 fp16 parameters,可直接 forward;ZeRO-3 需要在每个 layer 计算前 all-gather 当前 layer 的 parameter shards,临时形成完整 layer parameters。
  3. Forward 计算:每个 rank 用自己的 batch shard 计算 forward,并保存必要 activations 或 activation checkpoints。ZeRO-R 可以把 activation checkpoints 分片保存,必要时 offload 到 CPU。
  4. Forward 后释放参数:ZeRO-1/2 保留完整 fp16 parameters;ZeRO-3 释放当前 layer 的非本地 parameter shards,只留下本 rank 负责的 shard。
  5. Backward 前重建参数:ZeRO-3 在反向经过某个 layer 前,再次 all-gather 该 layer parameters,因为反向计算也需要对应参数;ZeRO-1/2 可直接使用完整参数。
  6. Backward 产生梯度:每个 rank 得到 local gradients。ZeRO-2/3 按 bucket 处理梯度,避免整模型梯度全部长期常驻。
  7. Gradient reduce-scatter:ZeRO-2/3 对 gradient bucket 做 reduce-scatter,把全局平均后的 gradient shard 直接发送给负责该参数片的 rank。接收方保留本地 shard,其他 rank 不长期保存该部分 gradient。
  8. Optimizer update:每个 rank 用自己的 gradient shard 更新自己的 fp32 master-weight shard、Adam momentum shard、Adam variance shard,并得到更新后的 parameter shard。
  9. 参数同步或保持分片:ZeRO-1/2 在 step 末尾 all-gather 更新后的 parameter shards,使下一步开始时每个 rank 都有完整 fp16 parameters;ZeRO-3 保持参数分片状态,下一步 forward 时再按 layer 临时 all-gather。
  10. Step 结束:optimizer states 始终按 shard 保存;gradients 在 update 后释放;activations 在 backward 后释放;parameters 在 ZeRO-3 中回到 shard-only 持久状态。

这个流程可以压缩成一条状态流:

param shardsall-gatherlayer paramsforward/backwardlocal gradsreduce-scattergrad shardsAdam updateupdated param shards. \text{param shards} \xrightarrow{\mathrm{all\text{-}gather}} \text{layer params} \xrightarrow{\mathrm{forward/backward}} \text{local grads} \xrightarrow{\mathrm{reduce\text{-}scatter}} \text{grad shards} \xrightarrow{\mathrm{Adam\ update}} \text{updated param shards}.

ZeRO-2 少了 per-layer parameter all-gather / release,因为它仍长期保存完整 fp16 parameters;ZeRO-3 把这一步也纳入动态状态流,所以 model-state memory 最低,通信量也更高。

11. ZeRO-2、ZeRO-3 与 FSDP 的工程选择

截至 2026 年,ZeRO stage 的工程选择主要按瓶颈划分,使用频率取决于模型规模、显存预算和通信拓扑:

  • 如果完整 fp16/bf16 parameters 能放进每张 GPU,主要压力来自 Adam optimizer states 和 gradients,ZeRO-2 往往是更实用的折中。它切 optimizer states 和 gradients,但保留完整 parameters 常驻,避免 ZeRO-3 / FSDP 那种每层 forward/backward 前后的 parameter all-gather / release。因此在模型能放下、追求吞吐和实现稳定性时,ZeRO-2 很常见。
  • 如果单卡放不下完整 parameters,或想把 batch/context/model size 推到更大,ZeRO-3 / FSDP FULL_SHARD 更自然。它们把 parameters、gradients、optimizer states 都分片,计算时临时 all-gather 当前模块参数,backward 后 reduce-scatter gradients,再用 sharded optimizer states 更新本地 parameter shard。
  • 在现代 LLM pretraining / post-training 系统里,ZeRO stage 常和 tensor parallelism、pipeline parallelism、expert parallelism、sequence/context parallelism、activation checkpointing、offload、Megatron distributed optimizer 等组合使用。实际选择经常是“模型是否能本地常驻参数 + 通信拓扑是否能承受参数 all-gather + 框架生态”共同决定。

ZeRO-3 和 PyTorch FSDP 的关系可以理解为同一类 full-shard data parallel 思想的两套工程实现。Hugging Face Accelerate 文档明确把 FSDP FULL_SHARD 映射到 DeepSpeed ZeRO stage 3;PyTorch FSDP2 教程也描述了相同的数据流:forward/backward 前 all-gather sharded parameters,backward 中 reduce-scatter gradients,optimizer 用 sharded gradients 更新 sharded parameters 和 optimizer states。

差异主要在工程生态:

  • DeepSpeed ZeRO-3 更偏 DeepSpeed 配置体系,配套 ZeRO-Offload、ZeRO-Infinity、CPU/NVMe offload、MiCS、DeepSpeed optimizer 和参数协调机制。
  • PyTorch FSDP / FSDP2 是 PyTorch 原生路线,和 torch.distributed、DTensor、state dict、auto wrap、PyTorch optimizer/compile 生态结合更自然。
  • 两者都要求处理 parameter gather 的生命周期。ZeRO-3 文档强调参数会在 forward/backward 需要时自动 collect/partition,并对外部参数访问提供手动协调机制;FSDP 通过 wrap unit 控制 gather/release 粒度。

所以一句话判断是:ZeRO-2 是“参数仍能常驻时的高性价比分片”;ZeRO-3 / FSDP 是“参数也必须分片时的 full-shard data parallel”。在显存够用时,ZeRO-2 往往更少通信、更容易跑快;在显存不够时,ZeRO-3 / FSDP 提供更强的模型规模上限。

12. 后续复验指标

  1. 对现代 FSDP / DeepSpeed ZeRO-3 / Megatron distributed optimizer,复验每步通信量、peak memory、overlap 效果和 fragmentation。
  2. 对 MoE 模型,分开记录 dense model states、expert states、EP all-to-all、activation memory、router/load imbalance。
  3. 对 RLHF/RLVR,单独记录 actor/ref/reward/critic 的 sharding 策略、weight sync、old logprob recompute、rollout engine 参数同步和 policy update peak memory。
  4. 对 long-context 训练,分开记录 model-state memory 和 activation/KV/context-parallel memory,避免把 ZeRO 的 model-state 优势误读成所有显存瓶颈都已消除。

主要启发

  • 大模型训练的第一性原理账本应从状态生命周期开始写:参数、梯度、optimizer states、activation checkpoints、temporary buffers、fragmentation 分别有什么大小、何时需要完整、何时可以释放。
  • DP 的高效率来自粗粒度 compute;MP 的显存优势来自状态分片。ZeRO 的关键是把这两者解耦,用动态通信让 DP 也获得状态分片。
  • 系统论文里的“可训练更大模型”要拆成三层:能否放进显存、能否保持吞吐、能否在合理时间内收敛。ZeRO 主要解决前两层,并明确承认 1T 训练时间仍需要更大算力。
  • 后续阅读 RL 系统时,看到 FSDP、ZeRO、Megatron distributed optimizer、optimizer state sharding,应回到本文的三阶段模型,判断当前系统到底切了 optimizer states、gradients、parameters 中的哪几类。
  • 内存优化常会移动瓶颈。ZeRO-DP 降低 model states 后,activation、buffer、fragmentation、communication overlap、batch size 和 convergence 都会成为下一层约束。

局限

  1. 实测集中在 dense GPT-like Transformer、V100/DGX-2 和 Megatron-LM 2019 baseline;对当前 MoE、long-context、FP8/BF16、H100/B200、NVLink/NVSwitch/IB 拓扑需要重新跑。
  2. 完整 Pos+g+pP_{os+g+p} 的 trillion 级训练在本文中主要是理论和通信分析,端到端实测实现聚焦 ZeRO-100B。
  3. 论文评估偏系统吞吐和最大模型规模,对最终模型质量、收敛速度、large batch 泛化边界没有系统展开。
  4. CPU activation offload 的收益条件较窄,强依赖 batch size 扩大带来的吞吐收益能否覆盖 PCIe/CPU 内存传输。
  5. ZeRO 的 memory sharding 不能单独解决 data pipeline、optimizer geometry、loss scaling、numerical determinism、checkpointing、fault tolerance 和 serving/training consistency。

跨论文关系

  • 2001.083612203.15556:scaling laws 说明为什么要增加参数/数据/计算;ZeRO 说明在 2020 年硬件上如何突破 optimizer state 和 gradient memory,使更大参数规模可训练。
  • 2309.14509 DeepSpeed Ulysses:作者上有 Samyam RajbhandariYuxiong He 直接重叠。ZeRO 处理 model states redundancy;Ulysses 处理 sequence activation 和 attention QKV/context layout。两者共同说明 DeepSpeed 系统线如何把 aggregate GPU memory 和 interconnect 用到不同训练状态上。
  • 2205.141352307.08691:ZeRO 处理 distributed training state memory;FlashAttention 处理 attention kernel IO 和 work partitioning。两者共同把 memory hierarchy 和 data movement 作为 LLM scaling 的核心变量。
  • 2606.04662:ZeRO 降低 Adam 类 optimizer states 的显存代价;Muon 论文解释 optimizer update direction 的曲率收益。前者是系统层 memory feasibility,后者是优化几何层效率。
  • 2409.19256verl 官方仓库:HybridFlow/VERL 编排 RLHF/RLVR 多模型 dataflow;ZeRO/FSDP 类分片决定 actor、critic、reference、reward model 的训练侧显存边界。
  • slime 官方仓库2602.15763:slime/GLM-5 把 Megatron training、SGLang rollout、async RL 和 weight sync 放进 production post-training stack;ZeRO 是理解其 optimizer/gradient/parameter sharding 的基础语言。
  • 2026-04-24:DeepSeek-V4 面向 million-token MoE、compressed attention、Muon 和 deterministic kernels;ZeRO 提供早期 dense 大模型训练中 optimizer/gradient/parameter state partitioning 的底层坐标。
  • 2606.04101:ZeRO 解决 data-parallel state redundancy;UltraEP 解决 MoE expert load imbalance 和 expert-state movement。二者都把“状态在哪里、何时移动、如何避免 straggler”作为 scaling 关键。
  • 2605.142202025-09-10:ZeRO 关注 trainer 侧 memory/communication;TIM 和 batch-invariant inference 提醒当 trainer 与 rollout engine 分离后,还要验证 logprob、kernel 和数值路径一致性。

Reference Intake Brief

Target

  • Intended target system: 新增 ZeRO / DeepSpeed training memory optimization 独立论文笔记;更新 content/utility/papers-index.md
  • Existing related assets: content/utility/papers-index.md2001.083612203.155562205.141352307.086912409.19256
  • Proposed form: 新建独立 Markdown 文档;更新当前收录,并在跨论文关系中补充训练内存优化主题。

Reusable Elements

  1. Mixed-precision Adam model-state memory formula:2Ψ+2Ψ+12Ψ=16Ψ2\Psi+2\Psi+12\Psi=16\Psi
  2. ZeRO-DP 三阶段内存公式:4Ψ+KΨ/Nd4\Psi+K\Psi/N_d2Ψ+14Ψ/Nd2\Psi+14\Psi/N_d16Ψ/Nd16\Psi/N_d
  3. 通信 volume 对照:DP 2Ψ2\PsiPos+gP_{os+g} 2Ψ2\PsiPos+g+pP_{os+g+p} 3Ψ3\Psi
  4. 系统分层语言:model-state memory、residual memory、communication volume、compute power gap。

Risks

  • Copyright/over-copying: 已用转述和公式化总结,避免长段复制论文原文。
  • Unsourced or unverifiable claims: 作者、版本、实验硬件、公式和结果来自 arXiv abstract 与 TeX source;后续 DeepSpeed/FSDP 演进只作为本地关系推论。
  • Tone/brand mismatch: 使用技术归档语气,避免营销式表述。
  • Safety/compliance issues: 本文为训练系统和内存优化论文,无直接安全滥用流程。
  • Overlap with existing assets: 当前档案有 FlashAttention、Muon、HybridFlow/VERL、slime 等系统节点,但缺 ZeRO 这种训练 memory sharding 基础节点。

Skipped

Material Reason
DeepSpeed 后续 ZeRO-Offload / ZeRO-Infinity 论文 属于后续独立系统演进,本笔记只记录 1910.02054。
当前 DeepSpeed 文档细节 当前任务分析论文;实时框架 API 后续可独立归档。

Recommendation

Decision: merge

Why: 本文是大模型训练系统的基础节点,补齐本地档案中 optimizer state / gradient / parameter sharding 的原始语言,并能连接 scaling laws、FlashAttention、Muon、HybridFlow/VERL、slime、GLM/DeepSeek 和 UltraEP 等后续材料。