大语言模型并行策略(八):零冗余优化 (ZeRO)
什么是 ZeRO
数据并行会将模型复制到多个 rank 上,每个 rank 处理不同的数据。此时,每个 rank 都拥有完整的参数,并在训练过程中维护完整的梯度和优化器状态,因此这些数据存在冗余。这些状态信息并非始终需要使用,例如模型参数只会在当前层的前向计算和反向传播时用到。零冗余优化(ZeRO, Zero Redundancy Optimizer)的核心思想是对模型参数、梯度和优化器状态进行分片。数据并行中的每个 rank 只持有一个分片;需要完整信息时,再从其他 rank 获取相应分片。通过这种方式,ZeRO 可以显著降低训练大模型时的内存占用。
ZeRO 的原理
在训练过程中,模型参数、梯度和优化器状态会占用大量的 GPU 内存。模型参数、梯度和优化器状态的内存占用量如下:
| 项目 | 数据类型 | 内存占用 |
|---|---|---|
| 权重 | BF16/FP16 | 2 B/参数 |
| 梯度 | BF16/FP16 | 2 B/参数 |
| 优化器状态 | FP32 | 12 B/参数 |
| 合计 | 16 B/参数 |
梯度的内存占用量与模型参数量直接相关,而优化器状态占用的内存最多。这是因为 Adam 优化器需要为每个参数维护一个高精度副本,以及动量和二阶矩两个状态;这三者通常都以 FP32 精度存储,以保证训练的稳定性和收敛性。Adam 优化器的状态占用量如下:
| 状态 | 含义 | 精度 | 内存占用 |
|---|---|---|---|
| 权重 | 模型参数的高精度副本 | FP32 | 4 B/参数 |
| 动量(momentum / first moment) | 梯度的指数移动平均 | FP32 | 4 B/参数 |
| 二阶矩(variance / second moment) | 梯度平方的指数移动平均 | FP32 | 4 B/参数 |
| 合计 | 12 B/参数 |
在数据并行中,每个 rank 使用不同的数据进行训练,各 rank 的梯度需要汇聚后再更新参数。因此,在整个训练过程中,各 rank 上的参数、梯度和优化器状态完全相同。为了降低内存消耗,ZeRO 将模型参数、梯度和优化器状态分片,并将这些状态分布到不同的 rank 上,每个 rank 只持有一个分片。参数、梯度和优化器状态都可以分片,但分片的状态越多,通信开销越大。
ZeRO 最早在 ZeRO: Memory Optimizations Toward Training Trillion Parameter Models 中提出,论文中提出了三层渐进分片方案:
- :分片优化器状态,参数和梯度不分片
- :分片优化器状态和梯度,参数不分片
- :分片优化器状态、梯度和参数
微软在 DeepSpeed 中实现了 ZeRO,为了方便配置,引入 1、2、3 来表示 ZeRO 的三种分片策略:
- ZeRO-1 =
- ZeRO-2 =
- ZeRO-3 =
后文中将使用 ZeRO-1、ZeRO-2 和 ZeRO-3 来表示 ZeRO 的三种分片策略。
设模型全部参数量为 ,数据并行的 rank 数为 ,每个参数对应的优化器状态占用内存为 ,则 ZeRO-1、ZeRO-2 和 ZeRO-3 的参数、梯度和优化器状态的内存占用量如下:

来自:https://arxiv.org/pdf/1910.02054
从上表可以看出,当数据并行度为 64 时,ZeRO-1、ZeRO-2 和 ZeRO-3 的内存占用量分别为 31.4 GB、16.6 GB 和 1.9 GB,分别是基线方案的 26%、14% 和 1.6%。
ZeRO-1: 切分优化器状态
ZeRO-1 将优化器状态切分到不同的 rank 上,每个 rank 只负责更新部分参数。反向传播结束后,每个 rank 从其他 rank 获取自己负责的参数对应的梯度,然后使用本地维护的优化器状态更新这些参数。参数更新完成后,再将最新参数同步给其他 rank。
下图描述了 ZeRO-1 的计算流程和通信模式:
重绘自:The Ultra-Scale Playbook: Training LLMs on GPU Clusters
ZeRO-1 的整个计算流程如下:
- 各个 rank 独立使用不同的数据进行前向计算和反向传播,计算出梯度。
- 由于每个 rank 只持有一个优化器状态分片,因此使用 Reduce-Scatter 将各分片对应的梯度汇聚到相应的 rank 上。
- 每个 rank 使用本地维护的优化器状态和汇聚后的梯度来更新自己负责的参数。
- 各 rank 使用 All-Gather 从其他 rank 获取最新参数,完成参数同步。
使用数据并行的情况下,每个 rank 都需要维护完整的优化器状态,内存占用量为 。使用 ZeRO-1 后,每个 rank 仅需要维护 的优化器状态,内存占用量大幅下降。
设 为整个模型参数的总字节数。使用数据并行时,反向传播后需要使用 All-Reduce 汇聚各 rank 的梯度,每个 rank 需要发送和接收的总数据量为 。使用 ZeRO-1 后,反向传播后使用 Reduce-Scatter 汇聚各分片的梯度;参数更新后使用 All-Gather 获取更新后的参数。每个 rank 的总通信量为 。
因此,从通信量来看,ZeRO-1 与数据并行相同,只是通信模式不同:数据并行使用 All-Reduce 汇总梯度;ZeRO-1 则先使用 Reduce-Scatter 汇总分片梯度,再使用 All-Gather 获取更新后的参数。
All-Reduce 的通信量
All-Reduce 底层的实现就是 Reduce-Scatter + All-Gather。设参与通信的 rank 数量为 ,需要发送的数据量为 ,则 Reduce-Scatter 的通信量为:
All-Gather 的通信量为:
整个 All-Reduce 的通信量为:
当 较大时,通信开销接近于 。
Reduce-Scatter 的原理
Reduce-Scatter 的工作原理是将待发送数据分为 个小块,每个 GPU 负责其中的一个小块,GPU 之间通过多次通信,将每个 GPU 的小块数据进行汇总,最终每个 GPU 都会得到自己负责的小块的汇总结果。
下面是 4 个 GPU 执行 Reduce-Scatter 的示意图:
首先 GPU 0 和 GPU 2、GPU 1 和 GPU 3 相互交换 1/2 的数据。然后 GPU 0 和 GPU 1、GPU 2 和 GPU 3 相互交换 1/4 的数据,最终每个 GPU 都会得到自己负责的小块的汇总结果。
设 为单个 GPU 持有的总数据量, 为 GPU 数量。在整个 Reduce-Scatter 通信过程中,每个 GPU 只需要发送和接收 字节的数据,因此通信量不会随 GPU 数量增加而线性增长。
ZeRO-2:切分优化器状态和梯度
设完整的梯度 被切分为 个分片,在 ZeRO-1 中使用 Reduce-Scatter 将各分片的梯度进行汇聚时,分片 会收集所有 rank 上的 ,执行跨 rank 归约。最终在 rank 上只有 是有效的,其他的分片 是未经规约的。
基于以上观察,ZeRO-2 在规约梯度后,会立刻释放掉非本地的梯度分片 ,从而降低内存占用。ZeRO-2 只需要保留 的梯度,因此 ZeRO-2 的内存占用量更低。
下面是 ZeRO-2 的计算流程:
ZeRO-2 的通信量与 ZeRO-1 相同,都是 。ZeRO-2 在 ZeRO-1 的基础上对梯度进行分片,并及时释放不再使用的梯度分片,从而降低内存占用量。从理论上看,ZeRO-2 相比 ZeRO-1 只有收益,ZeRO-1 似乎没有存在的必要。但在工程实践中,ZeRO-1 保存的完整梯度可用于调试,例如排查 NaN 异常值。此外,相比优化器状态,梯度占用的内存较小,因此对于参数量较小的模型,ZeRO-2 不会带来显著的内存节省。
对比原始论文中 和 的描述,ZeRO‑1 和 ZeRO‑2 这两者唯一的差异如下:
- ZeRO‑1:Reduce‑Scatter 结束后,不释放未归约梯度,持续占用显存;
- ZeRO‑2:Reduce‑Scatter 结束后,立即释放不属于 rank 的未归约梯度 ,降低显存占用。
ZeRO-3: 切分优化器状态、梯度和参数
ZeRO-3 在 ZeRO-2 的基础上进一步分片模型参数,每个 rank 仅保存 的参数。在执行前向计算时,需要获取完整参数;此时,每个 rank 使用 All-Gather 从其他 rank 获取完整参数,计算完成后,可以立刻释放非本地参数。
下面是 ZeRO-3 的前向计算流程:
在前向计算时,每个 rank 需要使用 All-Gather 获取完整参数,并在计算完成后立刻释放非本地参数。
下面是 ZeRO-3 的反向传播流程:
在反向传播中,每个 rank 同样需要使用 All-Gather 获取完整参数。
ZeRO-2 保留完整参数,而 ZeRO-3 对参数进行分片,每个 rank 只持有一个参数分片。ZeRO-2 在参数更新后使用 All-Gather 获取更新后的完整参数;ZeRO-3 则在需要参数时使用 All-Gather 获取完整参数,并在计算完成后立刻释放非本地参数。前向计算和反向传播都需要完整参数,因此 ZeRO-3 在一轮训练中需要两次使用 All-Gather 获取完整参数。相较于 ZeRO-2,ZeRO-3 的通信量多了一次 All-Gather,即 。
设模型参数总字节数为 ,则 ZeRO-1、ZeRO-2 和 ZeRO-3 的通信量如下:
| 策略 | 通信量 |
|---|---|
| DP | |
| ZeRO-1 | |
| ZeRO-2 | |
| ZeRO-3 |
总结
ZeRO 系列方法通过切分优化器状态、梯度和参数,有效降低单个 GPU 的内存占用。ZeRO-1 仅切分优化器状态,ZeRO-2 在此基础上切分梯度,ZeRO-3 进一步切分参数。尽管对这些状态进行了切分,ZeRO 系列方法仍保持与数据并行相同或相当的通信量。
为了提升训练速度而引入数据并行(DP)时,会产生显存冗余;ZeRO 系列方法的目的就是消除不同 DP 组之间的这部分冗余。
数据并行(DP)可以与张量并行(TP)、专家并行(EP)结合使用。例如,若有 16 张 GPU,可以划分为两个 DP 组,每个 DP 组使用 8 张 GPU 执行张量并行。此时两个 DP 组之间存在显存冗余,使用 ZeRO 系列方法可以消除这部分冗余。若只有 8 张 GPU 且 TP=8,则未引入数据并行,参数、梯度和优化器状态不存在 DP 造成的显存冗余,也就没有必要使用 ZeRO。