人人都会AI编程

9.5 ZeRO 优化:零冗余优化器的三级机制与显存优化原理

更新时间:2026-07-09

在 9.2 节中,我们介绍了最朴素的数据并行(Data Parallel, DP):每张 GPU 各自保存一份完整的模型副本,并行处理不同批次的数据,最后同步梯度。这种思路简单有效,却有一个致命的效率黑洞——显存冗余。如果你用 8 张卡做 DP,模型权重、梯度和优化器状态就在 8 张卡里重复存了 8 份。对于参数量动辄百亿的大模型,这种冗余直接把显存上限变成了模型规模的硬天花板。

ZeRO(Zero Redundancy Optimizer)正是为解决这个问题而生。它由微软 DeepSpeed 团队提出,核心思想极其直接:既然每张卡最终只更新自己那部分数据,那何必让所有卡都保存完整参数、梯度和优化器状态?

本节从显存拆解入手,逐层讲清 ZeRO 的三级机制、显存节省的数学原理、通信代价的隐藏技巧,以及它与张量并行(TP)、流水线并行(PP)的配合关系。


一、显存占用的“大头”到底在哪里?

在混合精度(FP16/BF16)训练中,Adam/AdamW 优化器对显存的消耗可以用一个公式粗略估算。假设模型参数量为 Ψ,则每参数占用显存为:

| 组件 | 精度 | 单参数字节 | 说明 |
|------|------|-----------|------|
| 模型参数(Weights) | FP16/BF16 | 2 | 前向/反向计算用 |
| 梯度(Gradients) | FP16/BF16 | 2 | 反向传播产出 |
| 优化器状态:FP32 Master Copy | FP32 | 4 | Adam 需要 32 位精度更新 |
| 优化器状态:Momentum | FP32 | 4 | 一阶动量 |
| 优化器状态:Variance | FP32 | 4 | 二阶动量 |
| 合计 | — | 16 字节 | 每参数 16 bytes |

这意味着一个 70B 参数的模型,仅模型状态(Parameter + Gradient + Optimizer States)就要吃掉 1.12 TB 显存。单张 A100 80GB 显存连零头都不够,即便 8 卡 DP,每张卡也要完整存 1.12 TB,显然不可能。

ZeRO 的三级优化,就是针对这 16 字节进行“分片手术”。


二、ZeRO 核心思想:分片(Partition)而非复制

ZeRO 没有改变 DP 的基本流程——数据仍被拆成不同微批次(micro-batch)送到各张卡,梯度仍需要全局聚合。改变的是模型状态的存储方式

  • 传统 DP:每张卡存完整的 16 Ψ 字节;
  • ZeRO:把参数、梯度、优化器状态切成 N 份(N 为数据并行度),每张卡只存 1/N,需要时通过集合通信实时聚合。

这带来了三级递进优化,分别对应 DeepSpeed 配置中的 stage=1/2/3


三、三级机制详解

ZeRO-1:优化器状态分片(Pos)

只对优化器状态(上述 12 字节/参数)进行切分。每张卡仅保存自己负责的那 1/N 参数的 FP32 Master Copy、Momentum 和 Variance。

  • 显存占用:每卡约 4Ψ + 12Ψ/N 字节。当 N 较大时趋近于 4Ψ,即显存降低约 4 倍
  • 通信量:与传统 DP 的 AllReduce 等价,可通过 Reduce-Scatter 实现,通信开销没有增加
  • 适用场景:显存瓶颈主要在优化器状态(Adam)时。7B–13B 模型的全参数微调常用。

ZeRO-2:梯度 + 优化器状态分片(Pos+g)

在 ZeRO-1 基础上,进一步对梯度(2 字节/参数)进行分片。反向传播时,各卡只计算并保留自己负责的 1/N 梯度,其余梯度在 Reduce 后直接释放。

  • 显存占用:每卡约 2Ψ + 14Ψ/N 字节。当 N 较大时趋近于 2Ψ,即显存降低约 8 倍
  • 通信量:仍与标准 DP 相当。梯度在反向传播时通过 Reduce-Scatter 直接落到对应卡上。
  • 适用场景:7B–30B 模型在 V100/A100 集群上的主流选择。是性价比最高的一级,显存省得多,通信没涨。

ZeRO-3:参数 + 梯度 + 优化器状态全分片(Pos+g+p)

最激进的一级,连参数(2 字节/参数)也进行分片。每张卡只存 1/N 的 FP16 参数。

  • 显存占用:每卡约 16Ψ/N 字节,模型状态显存与并行度 N 成线性反比。配合大 N,可以让模型状态的显存占用趋近于零(实际受激活值等其他开销限制)。
  • 通信量:前向传播需要 All-Gather 参数计算激活值,反向传播再次 All-Gather 参数并 Reduce-Scatter 梯度。总通信量是 ZeRO-1/2 的 1.5 倍
  • 适用场景:超大模型(70B+)或单节点卡数较多、节点内互联带宽充裕(NVLink)时。是“单卡装不下,但不想上 TP/PP”时的首选。

关键区别:ZeRO-3 虽然切分了参数,但它仍是数据并行范畴——每张卡仍承担全部层的计算,只是在计算前临时把参数凑齐。这与 9.3 节的张量并行(TP,把单层矩阵乘法拆开算)有本质不同。


四、通信代价与重叠:为什么 ZeRO 没有拖垮训练?

看到 ZeRO-3 多了 50% 通信,很多人的第一反应是“带宽会不会炸掉?”答案是:现代网络 + 通信重叠(Overlap)让代价基本可控。

  • ZeRO-1/2:梯度聚合本就发生在反向传播末尾,ZeRO 只是把它从 AllReduce 换成 Reduce-Scatter,通信量一致。在 InfiniBand 或 NVLink 环境下,这部分开销通常被计算掩盖。
  • ZeRO-3:参数 All-Gather 可以通过流水线与计算重叠。DeepSpeed 会提前把下一层需要的参数 All-Gather 好(prefetch),等计算层真正需要时,参数已经就位。此外,梯度 All-Gather 也可以与反向计算 overlap。

实用经验:在节点内 NVLink(带宽 600–900 GB/s)场景下,ZeRO-3 的通信开销通常只让整体训练速度下降 5%–15%;但如果是跨节点的慢速以太网(如 25–100 Gbps),ZeRO-3 的频繁参数聚合会成为瓶颈。此时应优先用 ZeRO-2,或配合 TP 减少 ZeRO 的跨节点压力。


五、ZeRO-Offload:显存不够,内存(甚至硬盘)来凑

如果 ZeRO-3 之后显存仍然不够,DeepSpeed 提供了 Offload 机制,把计算和存储进一步下沉:

| 级别 | 卸载目标 | 作用 | 代价 |
|------|----------|------|------|
| Offload Optimizer States | CPU 内存 | 把 Adam 的 FP32 状态、更新计算放到 CPU | 训练速度下降 10%–30%,但显存大幅释放 |
| Offload Parameters | CPU 内存 / NVMe SSD | 连 FP16 参数也放内存/硬盘 | 速度进一步下降,适合单机大模型微调 |
| ZeRO-Infinity | NVMe SSD 池 | 利用多机 SSD 聚合存储,支持万亿参数 | 极度依赖 PCIe 带宽和异步预取 |

真实案例:在单张 24GB RTX 3090 上微调 7B 模型,开启 ZeRO-3 + Offload Optimizer States 到 CPU 内存即可跑通;如果是 13B 模型,则需要 Offload Parameters 到 CPU,或配合梯度检查点(Gradient Checkpointing,用计算换显存)。


六、实战配置决策树

在实际工程中,如何选择 ZeRO 级别?

模型规模 ≤ 7B,单卡显存 ≥ 40GB?
  ├─ 是 → ZeRO-1 或直接用 DDP(无需 ZeRO)
  └─ 否 → 显存 24–32GB?
      ├─ 是 → ZeRO-2(性价比最高)
      └─ 否 → 显存 < 24GB?
          ├─ 是 → ZeRO-3 + Offload Optimizer
          └─ 更大模型(13B/30B/70B)→ ZeRO-3 + TP(单节点内)+ PP(跨节点)

DeepSpeed 配置片段示意

{
  "zero_optimization": {
    "stage": 2,
    "offload_optimizer": {
      "device": "cpu",
      "pin_memory": true
    },
    "allgather_partitions": true,
    "allgather_bucket_size": 2e8,
    "overlap_comm": true,
    "reduce_scatter": true
  },
  "gradient_checkpointing": true,
  "bf16": {"enabled": true}
}

关键调参项

  • overlap_comm: true:让梯度通信与反向计算重叠,减少延迟;
  • allgather_bucket_size / reduce_bucket_size:控制通信分块粒度,网络好时调大,网络差时调小;
  • contiguous_gradients:让梯度在内存中连续存储,减少碎片和拷贝开销。

七、ZeRO 与 TP、PP 的边界与协同

  • ZeRO 是 DP 的显存优化器,不改变模型结构,不拆分单层计算。它与 TP/PP 是正交关系。
  • 典型组合(引出 9.6 节 3D 并行)
  • 单节点 8×A100:TP=8(利用 NVLink)+ ZeRO-1/2。此时 DP=1,ZeRO 退化为纯显存优化。
  • 多节点 64×A100:TP=8(节点内)+ PP=4(跨节点层间)+ DP=2(数据并行)。ZeRO 可以在 DP 组内进一步分片模型状态,形成 ZeRO-3 + TP + PP 的 3D 混合并行。
  • 何时单用 ZeRO,何时必须上 TP/PP?
  • 如果单卡能装下单层(如 Transformer Block 的矩阵乘法),优先 ZeRO,因为实现简单、通信模式清晰;
  • 如果单层参数量就超过单卡显存(如早期 GPT-3 的超大 FFN 层),则必须用 TP 把单层拆开。

八、小结

ZeRO 是大模型训练工程中最具性价比的显存优化技术,它用通信换显存,在不增加模型开发复杂度的情况下,把数据并行的显存效率提升了 4–8 倍(ZeRO-1/2),甚至实现模型状态的线性扩展(ZeRO-3)。

核心 takeaway:

  • ZeRO-1:切优化器,省 4 倍显存,通信不变,Adam 训练必备;
  • ZeRO-2:加切梯度,省 8 倍显存,通信仍不变,最常用;
  • ZeRO-3:再切参数,显存线性降,通信涨 1.5 倍,适合大 DP 度或卡数充裕场景;
  • Offload:把显存压力转嫁给 CPU 内存/NVMe,单机魔改神器,但有速度损耗。

在 9.6 节中,我们将把 ZeRO 与 TP、PP 正式组合,讲解 3D 并行的协同策略——这是训练百亿到千亿参数模型的标准工程配方。