专栏大模型笔记·并行3 / 5
3 min学习

数据并行与 ZeRO:三级分片各省多少

数据并行是唯一不减少每卡显存的并行,ZeRO 用分片把它补上。从每个参数 16 字节的账开始,推 ZeRO-1、2、3 各自的每卡显存和通信量,再讲梯度分片放到 CPU 之后 GPU 上为什么只需要两个缓冲。

目录4 节

数据并行(DP)把 batch 切成 NN 份,每张卡跑一份完整的模型,每步结束把梯度 all-reduce 一次。它最简单,也最浪费:并行总览里,它是唯一一个不减少每卡显存的并行。ZeRO 就是为了把这一格补上。

每个参数 16 字节

先算普通 DP 下一个参数在一张卡上占多少字节。混合精度训练配 Adam 的经典账:

2BF16 权重+2BF16 梯度+4+4+4FP32 主权重、一阶矩、二阶矩=16 B.\underbrace{2}_{\text{BF16 权重}} + \underbrace{2}_{\text{BF16 梯度}} + \underbrace{4 + 4 + 4}_{\text{FP32 主权重、一阶矩、二阶矩}} = 16\ \text{B}.

后面 12 字节叫优化器状态,占了四分之三,而它们在整步训练里只在优化器 step 那一瞬间用到。换成 Muon 这类只有一阶动量的优化器,是 2+2+4+4=122 + 2 + 4 + 4 = 12 字节,优化器状态仍然占一半以上。梯度如果按 FP32 存,再加 2 字节。

一个 7B 模型按 16 字节算是 112 GB,一张 80 GB 的卡放不下。这就是 ZeRO 的出发点:NN 张卡上有 NN 份完全相同的状态,其中大部分每步只用一次,为什么不切开。

三级分片

切状态切梯度切权重普通 DP每卡一份完整的权重、梯度、优化器状态每参数 16 B(Adam)ZeRO-1优化器状态按ZeRO-2梯度也按ZeRO-3权重也按切,用时 all-gather,通信从涨到
ZeRO 的三级各多切一样东西。前两级不增加通信量,第三级为了在前向和反向时凑出完整权重要多一次 all-gather。 是参数量, 是数据并行的副本数。

ZeRO-1 切优化器状态。 每张卡只持有 1/N1/N 的 FP32 主权重和动量,只更新自己那一段参数,更新完把 BF16 权重 all-gather 回去。每参数变成 2+2+12/N2 + 2 + 12/N 字节。通信量不变:普通 DP 的 all-reduce 本来就等于一次 reduce-scatter 加一次 all-gather,ZeRO-1 只是把两次拆开,中间插进优化器 step。

ZeRO-2 再切梯度。 反向时每张卡算出的梯度做 reduce-scatter 而不是 all-reduce,每卡只留自己负责的那 1/N1/N,正好和它持有的优化器状态对上。每参数 2+14/N2 + 14/N。通信量还是 2Ψ2\PsiΨ\Psi 是参数量)。

ZeRO-3 连权重也切。 每张卡只持有 1/N1/N 的权重,前向和反向走到某一层时临时 all-gather 出完整的那层,用完就丢。每参数 16/N16/N,代价是多一次 all-gather,通信量变成 3Ψ3\Psi,并且每层都要等通信。

每参数字节(Adam)每步通信量代价
普通 DP16162Ψ2\Psi显存
ZeRO-14+12/N4 + 12/N2Ψ2\Psi
ZeRO-22+14/N2 + 14/N2Ψ2\Psi
ZeRO-316/N16/N3Ψ3\Psi每层多等一次 all-gather

前两级几乎是白拿的,所以现在几乎所有训练都至少开 ZeRO-1。第三级和张量并行、流水线并行在争同一件事(把权重切开),大模型训练里通常用后两者代替它,因为它们的通信更可控。

和别的并行叠在一起

ZeRO 里的 NN同一份权重有多少副本。叠上流水线并行和 expert 并行之后,这个数不再等于总卡数:一段流水线的权重只在同一个 stage 的卡之间复制,一个 expert 的权重只在同一个 EP rank 位置的卡之间复制。所以真正的分片因子是 DP 的度,而不是集群大小,估算显存时要用对。

在流水线并行下做梯度分片有一个额外的麻烦:反向是按 micro-batch 分批来的,梯度是累加出来的,不能等到最后再一次 reduce-scatter。做法是把梯度按模型 chunk 分批 reduce,边算边分片,Kimi K3 用的 Pipeline ZeRO-2 就是这一类。

分片放到 CPU:两个缓冲就够

ZeRO-2 之后每卡的梯度分片只有 1/N1/N,但它仍然是一块常驻显存。再进一步是把分片放到 CPU 内存里,GPU 上不再保存完整的梯度,只留一个double buffer:反向算出某个 chunk 的梯度进 buffer A,在 DP 之间 reduce,再累加到 CPU 上的分片;与此同时下一个 chunk 的梯度写进 buffer B。两个 buffer 轮流用,GPU 上的梯度占用从「全部参数」变成「两个 chunk」。

这是一个会反复出现的模式:任何「大块状态在别的地方、按块流过 GPU」的场景,GPU 上都只需要两个槽位,一个在用,一个在装。付的货币是 PCIe 带宽,条件是拷贝能藏在计算后面。Kimi K3 在训练里用它放梯度分片,在 RL 里又用同一对槽位把参考模型的权重流进来,见 K3 系统篇第 9 篇第 10 篇