ZERO训练技术

这篇笔记里我统一使用 ZeRO 这个写法,它指的是 DeepSpeed 提出的 Zero Redundancy Optimizer。理解 ZeRO 最好的方式,不是背“Stage 1/2/3 分别切了什么”,而是顺着一次训练 step 去看:

  1. 前向时每张卡手里有什么。
  2. 反向时梯度怎么产生、怎么通信、怎么留下来。
  3. 参数更新时谁来更、更新完怎么让所有卡重新对齐。

只要把这三件事看顺了,ZeRO-1ZeRO-2ZeRO-3 的差异就会非常清楚。


先立一个基线:普通 Data Parallel 怎么训练

在普通的数据并行 Data Parallel, DP 里,假设有 N 张卡,那么每一张卡都会保存完整的一份:

  • model parameters
  • gradients
  • optimizer states

这里的 optimizer states 以 Adam 为例,通常包括:

  • 参数本身对应的 FP16/BF16 权重
  • FP32 master weights
  • 一阶矩 m
  • 二阶矩 v

所以普通 DP 的问题不是“算不动”,而是“每张卡都重复保存一整套训练状态,显存浪费很大”。

一次训练 step 的流程

1. 前向

每张卡都拿着完整模型参数,各自处理不同的 micro-batch。
例如 rank 0 处理 batch 的一部分,rank 1 处理另一部分,但它们前向时用的模型权重是一模一样的完整副本。

2. 反向

每张卡先根据自己的 micro-batch 算出一套本地完整梯度
注意这时的梯度虽然只来自本卡数据,但张量形状是完整的,因为每张卡前向时持有的就是完整模型。

然后所有卡对这些完整梯度做一次 all-reduce

  • all-reduce 的作用是:各卡把同一个梯度张量做求和或求平均;
  • 操作结束后,每张卡都会拿到一份完全相同的完整梯度

3. 参数更新

由于每张卡现在都有:

  • 完整参数
  • 完整梯度
  • 完整 optimizer states

所以每张卡都会独立执行一遍完全一样的 Adam update。
因为输入完全一样,所以更新后的参数也完全一样,不需要额外同步。

普通 DP 的特点

  • 好处:逻辑最直接,实现简单。
  • 代价:参数、梯度、optimizer states 在每张卡上都完整复制,显存冗余最大。

后面的 ZeRO,本质上就是一步一步把这三类状态的“完整复制”改成“按 rank 分片保存”。


一张总对比表先建立全局感觉

方案 参数是否分片 梯度是否分片 Optimizer state 是否分片 前向是否需要 all-gather 反向主要通信 更新后如何恢复一致参数副本 通信开销趋势
普通 DP all-reduce 完整梯度 不需要额外恢复,所有卡本地更新结果天然一致 基线
ZeRO-1 all-reduce 完整梯度 更新后的参数分片通过 all-gather 重建完整参数 略高于 DP
ZeRO-2 layer-by-layer reduce-scatter 梯度 通常仍需 all-gather 让各卡重新拿到完整参数副本 高于 ZeRO-1
ZeRO-3 参数 all-gather + 梯度 reduce-scatter 参数平时就是分片存放,需要时再 all-gather 最高

如果只记一句话:

  • ZeRO-1 切 optimizer states。
  • ZeRO-2ZeRO-1 基础上再切 gradients。
  • ZeRO-3ZeRO-2 基础上再切 parameters。

下面按训练流程逐个展开。


ZeRO-1:只切 optimizer states

ZeRO-1 的核心是:模型参数和梯度仍然完整复制,但 optimizer states 不再在每张卡上都保存完整副本,而是按 rank 分片保存。

先回答最关键的一句

ZeRO-1 里:

  • 反向传播结束后,每张卡保留的梯度仍然是完整的
  • 被切分的重点是 optimizer states,而不是梯度;
  • 每张卡只更新自己负责的参数分片,是因为 optimizer states 的 ownership 已经按参数分片分配给不同 rank 了。

optimizer state 怎么切分

假设参数向量被逻辑上切成 N 份,对应 N 个 rank:

  • rank 0 负责第 0 段参数对应的 optimizer states
  • rank 1 负责第 1 段参数对应的 optimizer states
  • rank N-1 负责最后一段

这里“负责”指的是:

  • 这段参数的 m
  • 这段参数的 v
  • 这段参数的 FP32 master copy

只保存在对应的 rank 上,而不是每张卡都保存一份。

所以 ZeRO-1 节省的显存主要来自 Adam 状态,因为 Adam 的状态量通常比只看参数本身更大。

一次训练 step 的流程

1. 前向

和普通 DP 完全一样:

  • 每张卡都有完整参数副本;
  • 每张卡用自己的 micro-batch 做前向。

因此前向阶段并不需要额外参数通信。

2. 反向

和普通 DP 也基本一样:

  • 每张卡先得到一份本地完整梯度;
  • 然后对完整梯度做 all-reduce,得到所有卡一致的完整梯度。

所以到反向结束时,每张卡手里仍然有一整份完整梯度,而不是只保留自己的梯度分片。

3. 更新

这一步才是 ZeRO-1 和普通 DP 真正拉开差异的地方。

虽然每张卡都有完整梯度,但并不是每张卡都去更新整套参数,而是:

  1. 每个 rank 只更新自己负责的参数分片。
  2. 这次更新所需的 optimizer states 也只在这个 rank 上存在。
  3. 其他不归自己负责的参数段,本 rank 不做 optimizer update。

换句话说:

  • 梯度是完整可见的;
  • 但 optimizer states 是分片持有的;
  • 所以每个 rank 只对自己那一片参数拥有真正的“更新权限”。

更新后怎么得到统一模型

每个 rank 只更新了自己的参数分片之后,所有卡手里的参数副本其实暂时是不完整的“新状态”。
这时需要把每张卡更新好的参数分片重新汇总,恢复成一致的完整模型副本。

通常用的算子是 all-gather

all-gather 是什么

可以把 all-gather 理解成:

  • 每张卡拿出自己持有的一段张量;
  • 通信结束后,所有卡都拿到这些分片拼起来的完整张量。

所以在 ZeRO-1 里,all-gather 的作用就是:

  • rank 0 拿出自己更新后的参数分片;
  • rank 1 拿出自己更新后的参数分片;
  • 最后所有卡都重建出一份一致的完整参数副本。

all-gatherall-reduce 的区别

  • all-reduce 是“同位置元素做规约,比如求和/求平均,然后每张卡都得到同一个规约结果”。
  • all-gather 是“每张卡拿出不同分片,最后所有卡都拿到拼接后的完整结果”。

ZeRO-1 里:

  • 梯度同步靠 all-reduce
  • 更新后参数重建靠 all-gather

相比普通 DP,ZeRO-1 到底改了什么

相比普通 DP,ZeRO-1 唯一改变的是:
optimizer states 从“每卡完整保存”变成了“按 rank 分片保存”;参数和梯度依然是完整副本。


ZeRO-2:在 ZeRO-1 基础上再切 gradients

ZeRO-2 的核心是:不仅 optimizer states 分片,梯度也分片。
参数在大多数时刻仍然是完整副本,这一点和 ZeRO-1 一样。

最重要的区别是:ZeRO-2 不再让“整网完整梯度”在每张卡上一直保留到反向结束,而是尽量在反向过程中就把梯度规约并切走。

一次训练 step 的流程

1. 前向

前向和 ZeRO-1 一样:

  • 每张卡都持有完整参数副本;
  • 每张卡对自己的 micro-batch 做前向。

因此前向阶段仍然不需要额外参数 all-gather

2. 反向:重点是 layer-by-layer 的 reduce-scatter

这里是 ZeRO-2 的重点。

在普通 DP 或 ZeRO-1 里,一个直观想法是:

  • 整个模型先反向完;
  • 每张卡手里留着整网完整梯度;
  • 最后统一做梯度同步。

而在 ZeRO-2 里,更像是:

  1. 某一层一旦反向完成,这一层的梯度就已经产生了。
  2. 这时不等全模型其他层全部反向结束,就尽快对这一层梯度做 reduce-scatter
  3. 规约并切分完成后,本卡只保留自己负责的那一段梯度分片。
  4. 然后继续上一层的反向。

所以它是一个非常典型的 layer-by-layer 梯度规约与切分 的过程。

reduce-scatter 是什么

reduce-scatter 可以理解成两步的融合:

  1. 先做 reduce:把各卡对应梯度做求和或求平均。
  2. 再做 scatter:把规约后的完整结果按分片发给不同 rank。

因此它很像:

all-reduce + scatter

但作为一个融合算子,它避免了先在每张卡上留下完整规约结果、再手动切分出去的中间步骤。

反向后每张卡最终保留什么

这是 ZeRO-2 最重要的结论之一:

  • ZeRO-1 反向结束后,每张卡保留的是完整梯度
  • ZeRO-2 反向结束后,每张卡保留的是自己负责参数分片对应的那一部分已规约梯度

也就是说,在 ZeRO-2 里,完整梯度不会像 ZeRO-1 那样长期驻留在每张卡显存里。

这就是它比 ZeRO-1 更省显存的核心原因。

3. 更新

到了更新阶段,ZeRO-2 的逻辑反而比 ZeRO-1 更“对齐”:

  • optimizer states 本来就是按分片存的;
  • 梯度现在也已经按同样的 ownership 分片了;
  • 所以每个 rank 直接拿自己那部分梯度,更新自己那部分参数和 optimizer states。

这一步不再需要“虽然梯度完整,但我只更新其中一部分”的逻辑解释,因为梯度本身也已经被切到对应 rank 了。

4. 更新后参数如何一致

由于参数在训练大部分时候仍按完整副本形式参与前向,所以更新结束后,通常仍需要把各 rank 更新后的参数分片重新同步成一致的完整参数副本,以便下一轮前向继续直接使用完整参数。

这里可以继续理解为依赖参数分片的 all-gather 重建。

为什么 reduce-scatter 比“最后留完整梯度”更省显存

因为它改变了梯度在显存中的驻留方式:

  • 不是整网所有层的完整梯度一直堆在每张卡上;
  • 而是哪一层梯度算出来,就尽快规约并切成分片;
  • 本卡只留下自己需要负责更新的那一部分。

所以显存里长期保留的梯度体积明显下降。

相比 ZeRO-1,ZeRO-2 到底改了什么

相比 ZeRO-1ZeRO-2 新增的本质变化就是:
梯度不再完整复制,而是在反向过程中按 layer-by-layer 的 reduce-scatter 变成分片梯度。


ZeRO-3:在 ZeRO-2 基础上再切 parameters

ZeRO-3 的核心是:参数、梯度、optimizer states 三者全部分片。

这也是为什么它显存节省最激进,但通信也最重。

参数怎么切分

ZeRO-3 里,参数平时不是“每张卡都有完整副本”,而是:

  • rank 0 保存一部分参数
  • rank 1 保存另一部分参数
  • 每张卡只长期保存自己拥有的参数分片

所以 ZeRO-3ZeRO-2 最大的区别在于:
连前向所需的参数也不再常驻完整副本。

一次训练 step 的流程

1. 前向:按层 all-gather 参数

因为每张卡平时只有参数分片,而前向计算某一层时通常需要该层完整权重,所以流程会变成:

  1. 当前层要做前向之前,各 rank 先把这层参数分片通过 all-gather 拼成临时完整参数。
  2. 每张卡拿着这层完整参数做前向计算。
  3. 这一层算完后,非本 rank 所拥有的那部分参数可以释放掉,只保留本地参数分片。

所以 ZeRO-3 的前向不是“开局就有一整套完整模型”,而是“算哪一层,就临时 gather 哪一层”。

2. 反向:先需要完整参数,再做梯度 reduce-scatter

反向时同样要按层理解。

某一层反向时,需要该层参数参与梯度计算,因此通常也要保证这层的完整参数在计算时可见。
完成该层反向后,这层产生的梯度再像 ZeRO-2 一样走 reduce-scatter

  1. 该层反向得到梯度。
  2. 对这层梯度做 reduce-scatter
  3. 每张卡最终只保留自己负责参数分片对应的那部分梯度。

所以 ZeRO-3 的反向同时包含两类事情:

  • 为了算这一层,需要临时拿到这层完整参数;
  • 为了节省梯度显存,梯度出来后又尽快分片规约掉。

3. 更新

到了更新阶段,逻辑和 ZeRO-2 一脉相承:

  • 每个 rank 只持有自己的参数分片;
  • 只持有自己那部分梯度;
  • 只持有自己那部分 optimizer states;
  • 因此只更新自己的参数分片。

ZeRO-3 的通信量是不是变大

答案是:是的,通常会明显变大,而且会更频繁。

原因不是“梯度通信突然变得神秘地更多了”,而是参数本身也进入了按层通信流程:

  • ZeRO-2 主要增加的是梯度 reduce-scatter
  • ZeRO-3 则在前向和反向两边,都要频繁为当前层参数做 all-gather

这会带来两个后果:

  1. 通信从更粗粒度的 step 级,同步成了更细粒度的 layer 级。
  2. 参数 gather 和梯度 scatter 都要和计算过程紧密交织。

所以 ZeRO-3 通常对:

  • GPU 间互联带宽
  • 通信与计算重叠 overlap
  • bucket 化和 prefetch

会更加敏感。

为什么说 ZeRO-3 最省显存

因为三类主要训练状态都不再完整复制:

  • parameters 分片
  • gradients 分片
  • optimizer states 分片

这使得每张卡长期持有的 model states 最少。
代价就是:为了把“长期常驻显存”压到最低,必须接受更重的动态通信。

相比 ZeRO-2,ZeRO-3 到底改了什么

相比 ZeRO-2ZeRO-3 的本质新增变化就是:
参数也从“完整副本常驻”变成了“平时按分片保存,需要某层计算时再临时 all-gather”。


最后用一句话串起来

如果把三种 ZeRO 放成一条连续演进路线,可以这样记:

  1. ZeRO-1:参数和梯度还是完整的,先把 optimizer states 从复制改成分片。
  2. ZeRO-2:在 ZeRO-1 基础上,把梯度也从复制改成分片,关键算子是 layer-by-layer 的 reduce-scatter
  3. ZeRO-3:在 ZeRO-2 基础上,把参数也从复制改成分片,关键代价是前向和反向都要更频繁地做参数 all-gather

所以 ZeRO 的本质,不是某一个神奇算子,而是一套越来越激进的状态分片策略:

  • 先切 optimizer states
  • 再切 gradients
  • 最后切 parameters

显存越省,通信越重,这就是 ZeRO-1ZeRO-3 最核心的主线。


补充

这篇笔记先聚焦 ZeRO-1 / ZeRO-2 / ZeRO-3 的训练流程本身。
ZeRO-OffloadZeRO-Infinity、和 tensor parallel / pipeline parallel 的组合方式,这里先不展开,后面如果需要可以单独再开一篇。