数据并行与 ZeRO:三级分片各省多少
数据并行是唯一不减少每卡显存的并行,ZeRO 用分片把它补上。从每个参数 16 字节的账开始,推 ZeRO-1、2、3 各自的每卡显存和通信量,再讲梯度分片放到 CPU 之后 GPU 上为什么只需要两个缓冲。
数据并行(DP)把 batch 切成 份,每张卡跑一份完整的模型,每步结束把梯度 all-reduce 一次。它最简单,也最浪费:并行总览里,它是唯一一个不减少每卡显存的并行。ZeRO 就是为了把这一格补上。
每个参数 16 字节
先算普通 DP 下一个参数在一张卡上占多少字节。混合精度训练配 Adam 的经典账:
后面 12 字节叫优化器状态,占了四分之三,而它们在整步训练里只在优化器 step 那一瞬间用到。换成 Muon 这类只有一阶动量的优化器,是 字节,优化器状态仍然占一半以上。梯度如果按 FP32 存,再加 2 字节。
一个 7B 模型按 16 字节算是 112 GB,一张 80 GB 的卡放不下。这就是 ZeRO 的出发点: 张卡上有 份完全相同的状态,其中大部分每步只用一次,为什么不切开。
三级分片
ZeRO-1 切优化器状态。 每张卡只持有 的 FP32 主权重和动量,只更新自己那一段参数,更新完把 BF16 权重 all-gather 回去。每参数变成 字节。通信量不变:普通 DP 的 all-reduce 本来就等于一次 reduce-scatter 加一次 all-gather,ZeRO-1 只是把两次拆开,中间插进优化器 step。
ZeRO-2 再切梯度。 反向时每张卡算出的梯度做 reduce-scatter 而不是 all-reduce,每卡只留自己负责的那 ,正好和它持有的优化器状态对上。每参数 。通信量还是 ( 是参数量)。
ZeRO-3 连权重也切。 每张卡只持有 的权重,前向和反向走到某一层时临时 all-gather 出完整的那层,用完就丢。每参数 ,代价是多一次 all-gather,通信量变成 ,并且每层都要等通信。
| 每参数字节(Adam) | 每步通信量 | 代价 | |
|---|---|---|---|
| 普通 DP | 显存 | ||
| ZeRO-1 | 无 | ||
| ZeRO-2 | 无 | ||
| ZeRO-3 | 每层多等一次 all-gather |
前两级几乎是白拿的,所以现在几乎所有训练都至少开 ZeRO-1。第三级和张量并行、流水线并行在争同一件事(把权重切开),大模型训练里通常用后两者代替它,因为它们的通信更可控。
和别的并行叠在一起
ZeRO 里的 是同一份权重有多少副本。叠上流水线并行和 expert 并行之后,这个数不再等于总卡数:一段流水线的权重只在同一个 stage 的卡之间复制,一个 expert 的权重只在同一个 EP rank 位置的卡之间复制。所以真正的分片因子是 DP 的度,而不是集群大小,估算显存时要用对。
在流水线并行下做梯度分片有一个额外的麻烦:反向是按 micro-batch 分批来的,梯度是累加出来的,不能等到最后再一次 reduce-scatter。做法是把梯度按模型 chunk 分批 reduce,边算边分片,Kimi K3 用的 Pipeline ZeRO-2 就是这一类。
分片放到 CPU:两个缓冲就够
ZeRO-2 之后每卡的梯度分片只有 ,但它仍然是一块常驻显存。再进一步是把分片放到 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 篇。