ZeRO: Memory Optimizations Toward Training Trillion Parameter Models

基本信息

研究问题

大模型训练的显存瓶颈来自多个部分

论文从一个具体现象开始:1.5B 参数的 GPT-2 模型,16-bit 权重本身只需要约 3GB,但用普通 PyTorch 或 TensorFlow 训练时无法放入一张 32GB GPU。原因是训练时需要同时保存的不只是权重:

model states = parameters + gradients + optimizer states
residual states = activations + temporary buffers + fragmented memory

在 mixed-precision Adam 中,模型状态的每参数显存大致为:

  • fp16 parameters:2 bytes;
  • fp16 gradients:2 bytes;
  • fp32 master parameters:4 bytes;
  • fp32 first moment:4 bytes;
  • fp32 second moment:4 bytes。

因此,Adam 的模型状态合计约为 16 bytes/parameter。其中 optimizer-related 部分为 12 bytes/parameter,往往比实际用于 forward 的低精度权重更占显存。

activation 也可能成为主要瓶颈。论文以 1.5B GPT-2、sequence length 1024、batch size 32 为例,指出 activation 约需要 60GB;activation checkpointing 可以把它降到约 8GB,但对 100B GPT 类模型,即使 checkpointing 后,activation 仍可能约占 60GB。除此之外,大规模 gradient all-reduce、gradient norm 等操作可能需要额外的 flattened buffer;1.5B 参数模型的 fp32 flattened buffer 就可能需要约 6GB。内存碎片还可能导致“总剩余显存足够,但找不到连续空间”的 OOM。

Data Parallel 与 Model Parallel 的取舍

Data Parallel 的计算粒度大、通信量低、扩展性好,但每张 GPU 重复保存完整模型状态。Model Parallel 可以切分状态,显存效率更高,但会让计算粒度变小,并在 layer 内部或跨节点引入更多通信。

ZeRO 的目标是取得两者的组合特性:

保留 Data Parallel 的计算粒度和整体通信效率
        +
消除 Data Parallel 对模型状态的复制

其关键观察是:即使 parameters、gradients 和 optimizer states 被分片,训练过程也不要求每个 rank 在每个时刻都持有完整副本。某一层的参数只在该层 forward/backward 时需要,某个参数的梯度只需要被负责更新该参数分片的 rank 持有,optimizer state 更只在 optimizer step 时由对应 rank 使用。

核心主张

论文的主张可以压缩为以下几条:

  1. 只切分 optimizer states,就能在几乎不增加通信的情况下把模型状态显存降低约 4 倍。
  2. 再切分 gradients,可把模型状态显存降低约 8 倍,同时保持与标准 data parallel 接近的通信量。
  3. 再切分 parameters,模型状态显存可以随 data parallel degree 近似线性下降;完整 Pos+g+p 的通信量约为标准 data parallel 的 1.5 倍。
  4. activation checkpointing、临时 buffer 和 memory defragmentation 需要单独处理,不能由参数状态分片自动解决。
  5. ZeRO 与 model parallel 是互补关系。ZeRO 可以降低 data-parallel replicas 的冗余,model parallel 则继续处理单层计算、activation 或层深度带来的限制。
  6. 在论文的硬件和实现条件下,ZeRO-100B 能在 400 张 V100 上高效运行最高 170B 参数模型,并将 100B 规模训练吞吐提升到约 15 PFLOPs。

方法与机制

统一记号与普通 Data Parallel 基线

设:

  • :模型参数量;
  • :data parallel degree;
  • :optimizer states 的每参数 memory multiplier;
  • 、、:单个 parameter、gradient 和 optimizer-related state 的每参数字节数。

在论文的 mixed-precision Adam 例子中,、、,所以完整模型状态约为 bytes。标准 data parallel 在每个 rank 上复制全部模型状态:

这个公式只计算 model states,不包括 activation、temporary buffer、通信 workspace、allocator overhead 等 residual memory。因此实际运行时的峰值显存会高于这个下界。

Pos / ZeRO-1:Optimizer State Partitioning

将 optimizer states 平均分到 个 data-parallel ranks。第 个 rank 只保存并更新属于自己分片的 fp32 master weights、momentum 和 variance,而 parameters 与 gradients 仍在每个 rank 上完整复制。

每个 rank:
  full parameters
  full gradients
  1 / Nd optimizer states
 
optimizer step:
  each rank updates its own optimizer-state partition
  all-gather updated parameters

显存从:

下降为:

当 且 较大时,显存从约 降到约 ,对应约 4 倍的模型状态显存下降。因为每个 rank 仍需要完整参数和完整梯度,Pos 不能解决“参数本身就放不下一张 GPU”的问题。

通信方面,optimizer state partitioning 本身不要求额外的梯度通信;每个 rank 更新完自己负责的参数分片后,通过一次 all-gather 获取下一步所需的完整 updated parameters。论文将这一阶段的总通信量分析为与标准 data parallel 相同的量级。

Pos+g / ZeRO-2:Gradient Partitioning

在 Pos 的基础上,gradient 也按 parameter partition 分配。反向传播中某层 gradient 产生后,不再让所有 rank 都保留完整 gradient,而是只把对应部分 reduce 到负责该参数分片的 rank:

backward produces gradients
  -> bucket gradients by destination partition
  -> reduce each bucket to its owner rank
  -> release gradient copies after reduction
  -> owner rank applies optimizer update

这在 collective 语义上等价于 reduce-scatter:不同参数区间的 reduced gradient 最终落在不同 rank,而不是每个 rank 都拿到完整 reduced gradient。论文使用 bucketization,把同一目标 partition 的多个 gradient 合并后再通信,以获得更大的消息和更好的带宽,并让通信与 backward 计算重叠。

模型状态显存变为:

当 较大时,这一数值接近 ,相对于标准 data parallel 约为 8 倍降低。Pos+g 仍然要求每个 rank 持有完整 parameters,因此它主要解决 optimizer state 与 gradient 的冗余,不能单独突破参数模型本身的单卡容量限制。

Pos+g+p / ZeRO-3:Parameter Partitioning

第三阶段进一步按 data-parallel group 切分 parameters。每个 rank 长期只保留自己负责更新的 parameter partition;当 forward 或 backward 处理到其他 rank 所拥有的 partition 时,临时接收该部分参数,使用完后释放或重新分片。

论文使用动态通信调度来避免一次性 materialize 全部参数:

owner rank broadcasts the parameters for the current partition
  -> all ranks execute that partition
  -> parameters can be discarded after use
  -> continue with the next partition

模型状态显存近似变为:

这使模型状态显存随 近似线性下降。论文的理论表中,64-way DP 下,7.5B 模型的 model-state memory 从标准 DP 的 120GB 降到约 1.88GB;同样的分析给出 64-way DP 可容纳约 128B、1024-way DP 可容纳约 1T 参数模型的模型状态下界。

这里的“可容纳”只表示 model states 的容量分析,不等于完整训练一定可以运行。activation、通信临时 buffer、batch size、算力和训练时间仍然可能成为瓶颈。

三个阶段的通信账本

论文将标准 data parallel 的 gradient all-reduce 拆成两个阶段理解:reduce-scatter 与 all-gather。对大小为 的 gradient,pipelined all-reduce 的每个 rank 总 data movement 约为:

reduce-scatter: Ψ
all-gather:     Ψ
total:          2Ψ

在 Pos+g 中:

partitioned gradient reduce-scatter: Ψ
updated parameter all-gather:        Ψ
total:                               2Ψ

因此它在论文的带宽模型下与 baseline DP 相同。

在 Pos+g+p 中,参数需要在 forward 和 backward 各按需 all-gather 一次,梯度仍然需要 reduce-scatter:

parameter all-gather for forward:  Ψ
parameter all-gather for backward: Ψ
gradient reduce-scatter:            Ψ
total:                              3Ψ

总通信量约为 baseline 的 1.5 倍。这个结果依赖论文中的理想化带宽账本和有效的 pipeline/bucket 调度;实际 step time 还会受到消息大小、collective 实现、网络拓扑、通信与计算重叠程度以及参数访问粒度影响。

ZeRO-R:Residual Memory 优化

ZeRO-DP 只处理 model states。论文将 activation、temporary buffer 和 fragmentation 统称为 residual memory,并提出三类互补优化。

Partitioned Activation Checkpointing

在 model parallel 中,同一 activation 可能被多个 GPU 重复保存。Pa 与 activation checkpointing 配合:forward 后不保存完整 replicated activation,而是沿 model-parallel group 分片保存;backward 需要时再通过 all-gather 重建当前 activation。

forward:
  compute activation
  -> keep partitioned checkpoint
 
backward recomputation:
  all-gather the needed checkpoint
  -> recompute the cell/block
  -> run backward

论文给出的例子是:100B 模型、batch size 32、sequence length 1024、16-way model parallel 时,每层 checkpointed activation 约需要 33GB/GPU;分片后约降到 2GB/GPU。极大模型还可以将分片 checkpoint offload 到 CPU,形成 Pa+cpu,但代价是额外的数据搬运。

Constant-Size Buffers

大模型训练中,把全部 gradient 或中间结果融合成一个与模型大小成正比的 fp32 buffer,会形成新的显存瓶颈。例如 3B 参数模型的 32-bit fused buffer 约需要 12GB。ZeRO-R 使用与模型总参数量无关的 constant-size buffer,并在消息足够大与 buffer 显存占用之间取平衡。

Memory Defragmentation

activation checkpoint、重计算产生的短生命周期 tensor、parameter gradient 等长生命周期 tensor 交错分配,容易造成 memory fragmentation。ZeRO-R 预先分配连续的 activation checkpoint 和 gradient buffer,在运行时把新产生的 tensor 拷入这些区域,减少 allocator 搜索连续空间的开销和因碎片导致的 OOM。

ZeRO 与 Model Parallel 的组合

ZeRO 的分片维度是 data-parallel group,Model Parallel 的分片维度是模型计算本身。两者可以组合:

Tensor / Model Parallel:
  split a layer's computation and parameters
 
ZeRO-DP:
  split replicated model states across data-parallel replicas
 
Pipeline Parallel:
  split layers into stages when model depth or activation placement requires it

论文指出,单纯为了降低 model-state memory 时,ZeRO-DP 可能比跨节点 Model Parallel 更有效,因为它保留了较大的计算粒度和更低的跨节点通信压力。但 Model Parallel 仍有两类价值:它可以降低超大模型的 activation footprint,也可以在 data parallel alone 导致 global batch 过大时帮助控制 batch size。论文给出的组合分析中,-way data parallel 与 -way model parallel 的理论 model-state memory reduction 最多可达到 。

实验与证据

实现与硬件

论文在 PyTorch 中实现 ZeRO-100B,接口可以包装普通 torch.nn.Module,不要求用户修改模型定义;同时支持与 Megatron-LM 的 model parallel 组合。实验硬件是 400 张 NVIDIA V100,分布在 25 个 DGX-2 节点,节点间通信带宽为 800Gbps。

实验重点不是运行完整的 trillion-parameter 训练,而是在当时硬件上验证约 100B 规模的可运行性。因此,评测实现启用的是:

ZeRO-DP: Pos + g
ZeRO-R:  constant-size buffers + memory defragmentation
          + partitioned activation checkpointing where useful

论文将这一实现称为 ZeRO-100B。参数分片 Pp 在论文完整设计和理论分析中出现,但不是本轮 ZeRO-100B 评测的主要实现组件。

模型规模与吞吐

实验使用 GPT-2 类 Transformer,改变 hidden size、layer 数和 model-parallel degree,覆盖 1.5B 到 170B 参数规模。主要结果包括:

  • 与 Megatron-LM model-parallel baseline 相比,ZeRO-100B 在 400 GPU 上可以高效运行最高 170B 参数模型,论文称模型规模超过当时 SOTA 的 8 倍;
  • 对 8B 到 100B 模型,平均持续吞吐约为 15 PFLOPs,超过峰值算力的 30%;
  • 在大模型规模上,较 baseline 的训练速度最高提升约 10 倍;
  • 100B 模型在 400 张 V100 上达到每 GPU 超过 38 TFLOPs 的量级;
  • 使用 data parallel alone 时,ZeRO-100B 在 128 GPU 上支持最高约 13B 参数模型,而普通 DDP 在相同类型设置下约在 1.4B 左右就会遇到显存限制。

这些结果不能简单解释为“ZeRO 的通信更少”。在大模型实验中,显存下降允许每 GPU 放入更大的 batch size,从而提高 arithmetic intensity;同时,ZeRO 避免了大规模 model parallel 跨节点通信。因此吞吐收益来自显存、batch size、通信拓扑和计算粒度的共同作用。

Super-linear scalability

论文在 60B 模型上从 64 GPU 扩展到 400 GPU,观察到 super-linear speedup。原因不是增加 GPU 后每个 token 的基础计算量变少,而是 Pos+g 随 data-parallel degree 增大而降低每 GPU 的 model-state memory,使每 GPU 可以容纳更大的 batch size。batch 增大后,计算的 arithmetic intensity 提升,GPU 利用率和整体吞吐因此可能超过线性扩展。

这也是大模型系统中需要谨慎理解的“super-linear”:它依赖原先处在显存和 batch size 受限区间,而不是一种普遍适用的通信定律。若 batch size 已经足够大,或更大 batch 影响优化稳定性,额外 GPU 不一定继续产生相同收益。

配置消融与 ZeRO-R 的作用

论文用五组配置观察 residual memory 优化的效果:

配置ZeRO-DPZeRO-R
C1Posconstant-size buffers + memory defragmentation
C2PosC1 + partitioned activation checkpointing
C3Pos+gC1
C4Pos+gC2
C5Pos+gC2 + CPU offload of partitioned activations

在固定 batch size、16-way model parallel 的分析中,Pa 将可运行模型规模从约 40B 推到约 60B;加入 Pos+g 后进一步到约 140B;对 170B 模型,CPU offload 的 Pa+cpu 才能在极端显存约束下运行。CPU offload 并不总是提高吞吐:对 60B 模型,activation 在 GPU 与 CPU 间搬运的代价可能让性能下降;当模型已经无法运行或只能使用很小 batch 时,offload 才更有价值。

Turing-NLG

论文还报告了使用 ZeRO-100B 训练的 17B 参数 Turing-NLG。模型在 WebText-103 上的 validation perplexity 达到 10.21,并报告约 41.4 TFLOPs/GPU 的持续吞吐。这个结果说明 ZeRO 不只是容量分析,而是被用于一条完整语言模型训练路径;但该结果同时受到模型架构、数据、优化器和训练配方影响,不能把 perplexity 改善单独归因于 ZeRO。

关键结论

ZeRO 的本质是生命周期感知的状态分片

ZeRO 最重要的洞见不是“把 tensor 除以 GPU 数”,而是识别每类状态的使用生命周期:

  • optimizer states 只在参数更新时由 owner rank 使用;
  • parameter 只在对应 layer 的 forward/backward 期间需要;
  • gradient 在 reduce 到 owner rank 后即可释放其他副本;
  • activation checkpoint 只需在 backward 的重计算点恢复;
  • 大型 temporary buffer 不应随总模型大小无限增长。

因此,ZeRO 是“分片驻留 + 按需通信 + 及时释放”的运行时策略。显存节省和通信增加之间的平衡,取决于状态何时被访问以及通信是否可以和计算重叠。

ZeRO 不是只优化 optimizer

ZeRO 名称容易让人只想到 optimizer state,但论文的完整范围包括:

ZeRO-DP: optimizer states -> gradients -> parameters
ZeRO-R:  activations -> temporary buffers -> fragmentation

实际使用时需要先判断 OOM 来源。如果瓶颈是 Adam states,ZeRO-1 可能已经有效;如果 gradient 也占据大量显存,可考虑 ZeRO-2;如果完整 parameters 本身无法驻留,则需要 ZeRO-3/FSDP full shard 或其他 model-parallel 方案;如果主要问题是长序列 activation,仍需要 activation checkpointing、sequence/context parallel 或 activation offload。

论文中的理论容量不是训练吞吐承诺

论文表中的 128B、1T 等数字是 model-state memory 的容量推导。真正的训练系统还必须满足:

  • activation 和 temporary buffer 有足够空间;
  • 参数 all-gather、gradient reduce-scatter 能够被有效调度;
  • global batch size 不因扩展过大而影响优化;
  • checkpoint 能够保存和恢复分片状态;
  • 网络、CPU memory 和 GPU 数量足以支撑训练时间。

把“能放下”与“能高效训练”分开,是阅读分布式训练系统论文时非常重要的判断标准。

局限与疑问

  1. 实验与完整方案并不完全相同。 论文的理论方案包含 Pp,但主要的 100B 级实现是 Pos+g + ZeRO-R。阅读结果时不能把 ZeRO-3 的理论容量直接当成 ZeRO-100B 的实验结果。
  2. 通信模型做了理想化。 论文主要按带宽和 data movement 量分析通信,并没有覆盖所有网络 latency、collective 算法、消息粒度和拓扑差异;实际 performance 取决于实现细节。
  3. 显存公式主要针对 model states。 activation、allocator、CUDA kernel workspace、通信 buffer 和碎片会使真实峰值高于公式。
  4. 大 batch 带来的 super-linear speedup 有适用范围。 当训练已经超过 critical batch size,继续增大 batch 可能损害优化效率或收敛速度。
  5. CPU offload 是容量与吞吐的交换。 Pa+cpu 可以解决极端容量问题,但 CPU-GPU 数据搬运可能成为新的瓶颈。
  6. Checkpoint 和并行布局有工程耦合。 分片状态需要明确保存 owner、partition metadata、optimizer state 和 world-size 变化下的恢复策略;论文的容量分析本身不覆盖完整 checkpoint migration 问题。
  7. 与现代实现的名称不能机械等同。 今天的 DeepSpeed ZeRO、PyTorch FSDP、Megatron Distributed Optimizer 在具体参数聚合、flattening、prefetch、state dict 和 checkpoint 语义上都有演进,论文提供的是核心设计来源,不是当前 API 的完整说明。

分析与判断

这篇论文真正改变的是“data parallel 必须复制完整训练状态”这一默认前提。它并没有修改 language model 的 forward、backward 或 optimizer update,而是把训练状态的所有权从“每个 replica 都有一份”改成“每个 partition 有一个 owner,其他 rank 在需要时临时获得副本”。

从工程角度看,ZeRO 的价值可以分成三层:

  1. 容量层:把整个集群的 aggregate memory 变成可利用资源,让参数量不再直接受单卡 model-state memory 限制。
  2. 通信层:通过 reduce-scatter、all-gather、broadcast 和 bucketization,把分片状态重新组织成可执行的 forward/backward 流程。
  3. 系统层:再通过 activation partitioning、constant-size buffer 和 defragmentation,处理模型状态之外的真实显存问题。

因此,学习 ZeRO 时最值得掌握的不是三个 stage 的名称,而是下面这条推导链:

复制的状态是冗余的
  -> 按 data-parallel group 分片
  -> 根据状态生命周期按需 materialize
  -> 用 collective communication 保持计算语义
  -> 用 bucket / overlap / release 控制性能与峰值显存

这也解释了 ZeRO 为什么能与 Megatron 的 Tensor Parallel、Pipeline Parallel 组合:它们优化的是不同维度。Megatron 主要切分单层计算,ZeRO 主要消除 data-parallel replicas 的状态冗余;在超大规模训练中,二者通常不是二选一。

相关知识链接