ZERO训练技术
这篇笔记里我统一使用 ZeRO 这个写法,它指的是 DeepSpeed 提出的 Zero Redundancy Optimizer。理解 ZeRO 最好的方式,不是背“Stage 1/2/3 分别切了什么”,而是顺着一次训练 step 去看:
- 前向时每张卡手里有什么。
- 反向时梯度怎么产生、怎么通信、怎么留下来。
- 参数更新时谁来更、更新完怎么让所有卡重新对齐。
只要把这三件事看顺了,ZeRO-1、ZeRO-2、ZeRO-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-2在ZeRO-1基础上再切 gradients。ZeRO-3在ZeRO-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 真正拉开差异的地方。
虽然每张卡都有完整梯度,但并不是每张卡都去更新整套参数,而是:
- 每个 rank 只更新自己负责的参数分片。
- 这次更新所需的 optimizer states 也只在这个 rank 上存在。
- 其他不归自己负责的参数段,本 rank 不做 optimizer update。
换句话说:
- 梯度是完整可见的;
- 但 optimizer states 是分片持有的;
- 所以每个 rank 只对自己那一片参数拥有真正的“更新权限”。
更新后怎么得到统一模型
每个 rank 只更新了自己的参数分片之后,所有卡手里的参数副本其实暂时是不完整的“新状态”。
这时需要把每张卡更新好的参数分片重新汇总,恢复成一致的完整模型副本。
通常用的算子是 all-gather。
all-gather 是什么
可以把 all-gather 理解成:
- 每张卡拿出自己持有的一段张量;
- 通信结束后,所有卡都拿到这些分片拼起来的完整张量。
所以在 ZeRO-1 里,all-gather 的作用就是:
- rank 0 拿出自己更新后的参数分片;
- rank 1 拿出自己更新后的参数分片;
- …
- 最后所有卡都重建出一份一致的完整参数副本。
all-gather 和 all-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 里,更像是:
- 某一层一旦反向完成,这一层的梯度就已经产生了。
- 这时不等全模型其他层全部反向结束,就尽快对这一层梯度做
reduce-scatter。 - 规约并切分完成后,本卡只保留自己负责的那一段梯度分片。
- 然后继续上一层的反向。
所以它是一个非常典型的 layer-by-layer 梯度规约与切分 的过程。
reduce-scatter 是什么
reduce-scatter 可以理解成两步的融合:
- 先做 reduce:把各卡对应梯度做求和或求平均。
- 再做 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-1,ZeRO-2 新增的本质变化就是:
梯度不再完整复制,而是在反向过程中按 layer-by-layer 的 reduce-scatter 变成分片梯度。
ZeRO-3:在 ZeRO-2 基础上再切 parameters
ZeRO-3 的核心是:参数、梯度、optimizer states 三者全部分片。
这也是为什么它显存节省最激进,但通信也最重。
参数怎么切分
在 ZeRO-3 里,参数平时不是“每张卡都有完整副本”,而是:
- rank 0 保存一部分参数
- rank 1 保存另一部分参数
- …
- 每张卡只长期保存自己拥有的参数分片
所以 ZeRO-3 和 ZeRO-2 最大的区别在于:
连前向所需的参数也不再常驻完整副本。
一次训练 step 的流程
1. 前向:按层 all-gather 参数
因为每张卡平时只有参数分片,而前向计算某一层时通常需要该层完整权重,所以流程会变成:
- 当前层要做前向之前,各 rank 先把这层参数分片通过
all-gather拼成临时完整参数。 - 每张卡拿着这层完整参数做前向计算。
- 这一层算完后,非本 rank 所拥有的那部分参数可以释放掉,只保留本地参数分片。
所以 ZeRO-3 的前向不是“开局就有一整套完整模型”,而是“算哪一层,就临时 gather 哪一层”。
2. 反向:先需要完整参数,再做梯度 reduce-scatter
反向时同样要按层理解。
某一层反向时,需要该层参数参与梯度计算,因此通常也要保证这层的完整参数在计算时可见。
完成该层反向后,这层产生的梯度再像 ZeRO-2 一样走 reduce-scatter:
- 该层反向得到梯度。
- 对这层梯度做
reduce-scatter。 - 每张卡最终只保留自己负责参数分片对应的那部分梯度。
所以 ZeRO-3 的反向同时包含两类事情:
- 为了算这一层,需要临时拿到这层完整参数;
- 为了节省梯度显存,梯度出来后又尽快分片规约掉。
3. 更新
到了更新阶段,逻辑和 ZeRO-2 一脉相承:
- 每个 rank 只持有自己的参数分片;
- 只持有自己那部分梯度;
- 只持有自己那部分 optimizer states;
- 因此只更新自己的参数分片。
ZeRO-3 的通信量是不是变大
答案是:是的,通常会明显变大,而且会更频繁。
原因不是“梯度通信突然变得神秘地更多了”,而是参数本身也进入了按层通信流程:
ZeRO-2主要增加的是梯度reduce-scatter;ZeRO-3则在前向和反向两边,都要频繁为当前层参数做all-gather。
这会带来两个后果:
- 通信从更粗粒度的 step 级,同步成了更细粒度的 layer 级。
- 参数 gather 和梯度 scatter 都要和计算过程紧密交织。
所以 ZeRO-3 通常对:
- GPU 间互联带宽
- 通信与计算重叠
overlap - bucket 化和 prefetch
会更加敏感。
为什么说 ZeRO-3 最省显存
因为三类主要训练状态都不再完整复制:
- parameters 分片
- gradients 分片
- optimizer states 分片
这使得每张卡长期持有的 model states 最少。
代价就是:为了把“长期常驻显存”压到最低,必须接受更重的动态通信。
相比 ZeRO-2,ZeRO-3 到底改了什么
相比 ZeRO-2,ZeRO-3 的本质新增变化就是:
参数也从“完整副本常驻”变成了“平时按分片保存,需要某层计算时再临时 all-gather”。
最后用一句话串起来
如果把三种 ZeRO 放成一条连续演进路线,可以这样记:
ZeRO-1:参数和梯度还是完整的,先把 optimizer states 从复制改成分片。ZeRO-2:在ZeRO-1基础上,把梯度也从复制改成分片,关键算子是 layer-by-layer 的reduce-scatter。ZeRO-3:在ZeRO-2基础上,把参数也从复制改成分片,关键代价是前向和反向都要更频繁地做参数all-gather。
所以 ZeRO 的本质,不是某一个神奇算子,而是一套越来越激进的状态分片策略:
- 先切 optimizer states
- 再切 gradients
- 最后切 parameters
显存越省,通信越重,这就是 ZeRO-1 到 ZeRO-3 最核心的主线。
补充
这篇笔记先聚焦 ZeRO-1 / ZeRO-2 / ZeRO-3 的训练流程本身。
像 ZeRO-Offload、ZeRO-Infinity、和 tensor parallel / pipeline parallel 的组合方式,这里先不展开,后面如果需要可以单独再开一篇。