专栏大模型笔记·性能优化4 / 5
4 min学习

显存账本:每张卡上住着什么,搬走付什么

「性能优化」章的开篇。训练时一张卡的显存分五类:权重、梯度、优化器状态、激活、通信缓冲。前三类由参数量和并行度决定,激活由 micro-batch 长度和流水线在途份数决定。每种省显存的技巧都是把账本某一行搬走并付一种货币:FLOPs、精度、PCIe、网络。

目录6 节

显存不够是训练大模型时最先撞上的墙。撞墙之后能做的事很多,重算、量化、offload、换并行,但要选对,得先知道显存到底被谁占了。这一篇给一个通用的记账方法,后面的文章都在这张账本上做加减。

五类

训练时一张卡上的显存分五类:

类别由什么决定谁能把它变小
权重参数量 ÷ 切它的并行度TP、PP、EP、ZeRO-3
梯度与权重同形状同上,加 ZeRO-2
优化器状态每参数 8 到 12 字节ZeRO-1
激活micro-batch 长度 × 每 token 每层保存量 × 在途份数重算、量化、offload、减小 micro-batch
通信缓冲all-to-all、all-gather 的临时空间固定 shape、复用

前三类是静态的,模型和并行度定了它们就定了,ZeRO 那篇算过每参数的字节数。激活是唯一随输入长度增长的一项,也是优化空间最大的一项。

激活这一行怎么算

反向传播需要前向的中间结果。每一层为反向保存的张量,按「每个 token 多少个值」数一遍:注意力层是输入、Q、K、V、注意力输出、门之类,通常是隐藏维的几倍;MoE 层是每个 token 激活的 KK 个 expert 各自的中间激活,通常是最大的一行。把所有层加起来得到每 token 的总保存量 aa,再乘三个数:

激活显存=a×S×1P×n在途×每值字节数.\text{激活显存} = a \times S \times \frac{1}{P} \times n_{\text{在途}} \times \text{每值字节数}.

SS 是 micro-batch 的 token 数,1/P1/P 是这张卡只持有 1/P1/P 的层,n在途n_{\text{在途}} 是流水线里同时攒着几份 micro-batch 的激活。最后这个数在 1F1B 下是 PrP - r(rank 0 最多,见流水线并行),虚拟段会让它更大。所以激活显存的形状是一个阶梯:前面的 rank 满,后面的 rank 空。

一个直觉性的数量级:一个几千亿到几万亿参数的 MoE 模型,每 token 的总保存量在千万个值这个量级,BF16 下几十 MB。乘上几千的 SS 和几份在途,不做任何优化就是几百 GB,远超一张卡。

搬走一行,付一种货币

激活相关的每个技巧都是把账本某一行搬到别处,然后付一种货币。三种基本策略:

策略省的付的适合谁
重算整个张量的显存反向时再做一次前向的 FLOPs算得快、占得大的逐元素算子(激活函数、归一化)
量化一半到四分之三的字节精度对精度不敏感的中间激活,配 block-wise 的缩放
offload整个张量的显存PCIe 带宽(到 CPU)或网络带宽(到别的卡)反向时才用、间隔长的张量

三种可以叠:先量化再 offload,逐元素的重算。现代的做法是把它们抽象成同一个接口,每个张量绑一种「存储策略」,用注解声明,与模型代码解耦,Kimi K3 的统一激活管理器就是这样做的。

offload 的条件是搬运能藏在计算后面。粗算:每张卡每步产生的激活字节数除以这一步的计算时间,得到需要的带宽;PCIe 5 单向约 64 GB/s,NVLink 几百 GB/s,跨机网络几十 GB/s。带宽不够的时候 offload 会让计算等搬运,这时只能退回重算或者减小 micro-batch。

激活的阶梯形状还给了第四种办法:在流水线 rank 之间搬。后面的 rank 有空闲显存,前面的 rank 满,把前面 rank 的激活远程 offload 到后面 rank 的显存里,各 rank 拉平。这付的是卡间网络的带宽,比 CPU offload 快。

静态三行也能动

权重、梯度、优化器状态这三行通常靠并行度决定,但也有搬的余地。梯度分片可以放到 CPU,GPU 上只留两个 chunk 大小的 double buffer;优化器 step 可以在 CPU 上做(ZeRO-Offload 的思路);某些优化器需要的临时完整矩阵(比如 Muon 的正交化)可以改成按需拉取而不是整份 all-gather。这些都在 ZeRO 那篇和具体模型的案例里。

通信缓冲

最容易被忘的一行。MoE 的 all-to-all 要预留接收缓冲,如果每个 rank 会收到多少 token 事先不知道,就得按最坏情况留,最坏情况是所有 token 涌向同一个 rank,缓冲要 RR 倍。任何能让通信 shape 静态化的设计(比如让每个 rank 恰好收到相同份数的 token)都能把这一行缩到固定大小。

用这张账本

面对一个放不下的模型,按行问:

  1. 静态三行是多少?能不能靠 ZeRO 和并行度压到目标以下?
  2. 激活是多少?哪一层最大?它的那部分能重算、能量化、还是只能 offload?
  3. 流水线的阶梯有多陡?前面的 rank 有没有办法借后面 rank 的显存?
  4. 通信缓冲是不是按最坏情况留的?

一个把这四步走完的真实例子是 Kimi K3 的显存篇,那篇有一张可以拖并行度的账本。