显存不够用,是所有训练大模型的人都得过的病。

我第一次正经训大模型,是拿单卡跑一个几十亿参数的模型,loss 还没开始降,显存先爆了,CUDA out of memory 刷了满屏。那会儿我唯一的念头是:这东西到底是谁在吃显存。

罪魁是三件套:模型参数 P、梯度 G、优化器状态 S。参数是模型本身的体重;梯度是每一轮算出来的"该往哪走";优化器状态(动量、二阶矩那些)是 Adam 这类优化器自己记的账本。传统训练里,每张卡这三样都得存一份,显存等于 $P + G + S$ 各来一遍。

DeepSpeed 的 ZeRO(Zero Redundancy Optimizer)干的事,一句话:把这三样从"每卡都存一份"改成"大家分着存,用的时候再凑齐"。它分三级,一级比一级狠。

Stage 1:先分优化器状态

优化器状态通常是最肥的。一个 Adam 就要给每个参数额外记两份浮点(动量 + 二阶矩),再加上 fp32 主权重,S 往往是 P 的好几倍。Stage 1 只分它:S 摊到 N 张卡上,每卡显存变成 $P + G + \frac{S}{N}$。

省,但省得不够狠。

Stage 2:梯度也分了

Stage 2 把梯度 G 也摊出去,每卡 $P + \frac{G}{N} + \frac{S}{N}$。注意,参数 P 每张卡还留着一整份。

这是当年最常用的档位。理由很实在:参数不跨卡,前向、反向都不用为"取参数"额外通信,开销小,速度快。显存省了,训练节奏基本不变。

Stage 3:参数也分了

Stage 3 才动真格的,参数 P 也摊成 $\frac{P}{N}$,每卡只剩 $\frac{P + G + S}{N}$,显存需求直接砍到接近 $\frac{1}{N}$。70B 这种在 Stage 2 下要好几张卡才塞得下的大家伙,Stage 3 里单卡也能跑,只是慢。

代价在通信。参数分散在别人那儿,每一步前向要 gather 一次、反向算完梯度还要 all-reduce 一次,通信量比 Stage 2 多出一截。高带宽集群(NVLink、InfiniBand 互联)里这钱花得不心疼;千兆以太网那种环境,Stage 3 的墙钟时间可能比 Stage 2 还难看。

所以 Stage 2 和 Stage 3 之间,压根没有谁碾压谁:省显存是 Stage 3 赢,快速度是 Stage 2 赢,带宽决定 Stage 3 的账算不算得过来。

对模型精度有没有影响

没有。这三档都是把同一份训练在硬件上的摆法换了个花样,每一步算出来的数值在浮点误差范围内完全一致,模型最终效果一样。它们换的只是显存和速度,不是模型本身。

怎么选

三板斧:

  1. 先看模型多大。几十亿参数、显存够呛,Stage 2 基本够用;几百亿往上的,只能 Stage 3。
  2. 再看带宽。机内 NVLink、跨机 InfiniBand,放心上 Stage 3;普通千兆网,Stage 2 更稳妥。
  3. 最后看目标。显存是瓶颈,Stage 3;训练时间紧张、显存还撑得住,Stage 2。

实践里还有个不花钱的加分项:混合精度。fp16/bf16 能把显存再压一半,配合 ZeRO 效果叠加,是训练标配。

ZeRO 的三档,是显存和通信之间的三档交易。选哪档,取决于你的模型、你的网线,和你对训练时间的耐心。

FIN