WangYu::Space

cat /dev/mind

大语言模型并行策略(八):零冗余优化 (ZeRO)

分类:机器学习标签: LLM创建时间:2026-06-27 22:10:00

什么是 ZeRO

数据并行会将模型复制到多个 rank 上,每个 rank 处理不同的数据。此时,每个 rank 都拥有完整的参数,并在训练过程中维护完整的梯度和优化器状态,因此这些数据存在冗余。这些状态信息并非始终需要使用,例如模型参数只会在当前层的前向计算和反向传播时用到。零冗余优化(ZeRO, Zero Redundancy Optimizer)的核心思想是对模型参数、梯度和优化器状态进行分片。数据并行中的每个 rank 只持有一个分片;需要完整信息时,再从其他 rank 获取相应分片。通过这种方式,ZeRO 可以显著降低训练大模型时的内存占用。

ZeRO 的原理

在训练过程中,模型参数、梯度和优化器状态会占用大量的 GPU 内存。模型参数、梯度和优化器状态的内存占用量如下:

项目数据类型内存占用
权重BF16/FP162 B/参数
梯度BF16/FP162 B/参数
优化器状态FP3212 B/参数
合计16 B/参数

梯度的内存占用量与模型参数量直接相关,而优化器状态占用的内存最多。这是因为 Adam 优化器需要为每个参数维护一个高精度副本,以及动量和二阶矩两个状态;这三者通常都以 FP32 精度存储,以保证训练的稳定性和收敛性。Adam 优化器的状态占用量如下:

状态含义精度内存占用
权重模型参数的高精度副本FP324 B/参数
动量(momentum / first moment)梯度的指数移动平均FP324 B/参数
二阶矩(variance / second moment)梯度平方的指数移动平均FP324 B/参数
合计12 B/参数

在数据并行中,每个 rank 使用不同的数据进行训练,各 rank 的梯度需要汇聚后再更新参数。因此,在整个训练过程中,各 rank 上的参数、梯度和优化器状态完全相同。为了降低内存消耗,ZeRO 将模型参数、梯度和优化器状态分片,并将这些状态分布到不同的 rank 上,每个 rank 只持有一个分片。参数、梯度和优化器状态都可以分片,但分片的状态越多,通信开销越大。

ZeRO 最早在 ZeRO: Memory Optimizations Toward Training Trillion Parameter Models 中提出,论文中提出了三层渐进分片方案:

  1. PoP_o​:分片优化器状态,参数和梯度不分片
  2. Po+gP_{o+g}:分片优化器状态和梯度,参数不分片
  3. Po+g+pP_{o+g+p}:分片优化器状态、梯度和参数

微软在 DeepSpeed 中实现了 ZeRO,为了方便配置,引入 1、2、3 来表示 ZeRO 的三种分片策略:

后文中将使用 ZeRO-1、ZeRO-2 和 ZeRO-3 来表示 ZeRO 的三种分片策略。

设模型全部参数量为 Ψ\Psi,数据并行的 rank 数为 NdN_d,每个参数对应的优化器状态占用内存为 12B12B,则 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 的整个计算流程如下:

  1. 各个 rank 独立使用不同的数据进行前向计算和反向传播,计算出梯度。
  2. 由于每个 rank 只持有一个优化器状态分片,因此使用 Reduce-Scatter 将各分片对应的梯度汇聚到相应的 rank 上。
  3. 每个 rank 使用本地维护的优化器状态和汇聚后的梯度来更新自己负责的参数。
  4. 各 rank 使用 All-Gather 从其他 rank 获取最新参数,完成参数同步。

使用数据并行的情况下,每个 rank 都需要维护完整的优化器状态,内存占用量为 12Ψ12\Psi。使用 ZeRO-1 后,每个 rank 仅需要维护 1Nd\frac{1}{N_d} 的优化器状态,内存占用量大幅下降。

Φ\Phi 为整个模型参数的总字节数。使用数据并行时,反向传播后需要使用 All-Reduce 汇聚各 rank 的梯度,每个 rank 需要发送和接收的总数据量为 2Φ2\Phi。使用 ZeRO-1 后,反向传播后使用 Reduce-Scatter 汇聚各分片的梯度;参数更新后使用 All-Gather 获取更新后的参数。每个 rank 的总通信量为 Φ+Φ\Phi + \Phi

因此,从通信量来看,ZeRO-1 与数据并行相同,只是通信模式不同:数据并行使用 All-Reduce 汇总梯度;ZeRO-1 则先使用 Reduce-Scatter 汇总分片梯度,再使用 All-Gather 获取更新后的参数。

All-Reduce 的通信量

All-Reduce 底层的实现就是 Reduce-Scatter + All-Gather。设参与通信的 rank 数量为 NdN_d,需要发送的数据量为 MM,则 Reduce-Scatter 的通信量为:

M×Nd1NdM \times \frac{N_d-1}{N_d}

All-Gather 的通信量为:

M×Nd1NdM \times \frac{N_d-1}{N_d}

整个 All-Reduce 的通信量为:

2×M×Nd1Nd2 \times M \times \frac{N_d-1}{N_d}

NdN_d 较大时,通信开销接近于 2M2M

Reduce-Scatter 的原理

Reduce-Scatter 的工作原理是将待发送数据分为 NdN_d 个小块,每个 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 都会得到自己负责的小块的汇总结果。

MM 为单个 GPU 持有的总数据量,NdN_d 为 GPU 数量。在整个 Reduce-Scatter 通信过程中,每个 GPU 只需要发送和接收 M×Nd1NdM \times \frac{N_d-1}{N_d} 字节的数据,因此通信量不会随 GPU 数量增加而线性增长。

ZeRO-2:切分优化器状态和梯度

设完整的梯度 GG 被切分为 NdN_d 个分片,在 ZeRO-1 中使用 Reduce-Scatter 将各分片的梯度进行汇聚时,分片 ii 会收集所有 rank 上的 G[i]G[i],执行跨 rank 归约。最终在 rank rr 上只有 G[r]G[r] 是有效的,其他的分片 G[ir]G[i \neq r] 是未经规约的。

基于以上观察,ZeRO-2 在规约梯度后,会立刻释放掉非本地的梯度分片 G[ir]G[i \neq r],从而降低内存占用。ZeRO-2 只需要保留 1Nd\frac{1}{N_d} 的梯度,因此 ZeRO-2 的内存占用量更低。

下面是 ZeRO-2 的计算流程:

ZeRO-2 的通信量与 ZeRO-1 相同,都是 2Φ2\Phi。ZeRO-2 在 ZeRO-1 的基础上对梯度进行分片,并及时释放不再使用的梯度分片,从而降低内存占用量。从理论上看,ZeRO-2 相比 ZeRO-1 只有收益,ZeRO-1 似乎没有存在的必要。但在工程实践中,ZeRO-1 保存的完整梯度可用于调试,例如排查 NaN 异常值。此外,相比优化器状态,梯度占用的内存较小,因此对于参数量较小的模型,ZeRO-2 不会带来显著的内存节省。

对比原始论文中 PoP_oPo+gP_{o+g} 的描述,ZeRO‑1 和 ZeRO‑2 这两者唯一的差异如下:

ZeRO-3: 切分优化器状态、梯度和参数

ZeRO-3 在 ZeRO-2 的基础上进一步分片模型参数,每个 rank 仅保存 1Nd\frac{1}{N_d} 的参数。在执行前向计算时,需要获取完整参数;此时,每个 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,即 Ψ\Psi

设模型参数总字节数为 Ψ\Psi,则 ZeRO-1、ZeRO-2 和 ZeRO-3 的通信量如下:

策略通信量
DP2Ψ2\Psi
ZeRO-12Ψ2\Psi
ZeRO-22Ψ2\Psi
ZeRO-33Ψ3\Psi

总结

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。

评论 评论内容仅博主可见,不会公开显示)