显存账本:每张卡上住着什么,搬走付什么
「性能优化」章的开篇。训练时一张卡的显存分五类:权重、梯度、优化器状态、激活、通信缓冲。前三类由参数量和并行度决定,激活由 micro-batch 长度和流水线在途份数决定。每种省显存的技巧都是把账本某一行搬走并付一种货币:FLOPs、精度、PCIe、网络。
显存不够是训练大模型时最先撞上的墙。撞墙之后能做的事很多,重算、量化、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 激活的 个 expert 各自的中间激活,通常是最大的一行。把所有层加起来得到每 token 的总保存量 ,再乘三个数:
是 micro-batch 的 token 数, 是这张卡只持有 的层, 是流水线里同时攒着几份 micro-batch 的激活。最后这个数在 1F1B 下是 (rank 0 最多,见流水线并行),虚拟段会让它更大。所以激活显存的形状是一个阶梯:前面的 rank 满,后面的 rank 空。
一个直觉性的数量级:一个几千亿到几万亿参数的 MoE 模型,每 token 的总保存量在千万个值这个量级,BF16 下几十 MB。乘上几千的 和几份在途,不做任何优化就是几百 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,缓冲要 倍。任何能让通信 shape 静态化的设计(比如让每个 rank 恰好收到相同份数的 token)都能把这一行缩到固定大小。
用这张账本
面对一个放不下的模型,按行问:
- 静态三行是多少?能不能靠 ZeRO 和并行度压到目标以下?
- 激活是多少?哪一层最大?它的那部分能重算、能量化、还是只能 offload?
- 流水线的阶梯有多陡?前面的 rank 有没有办法借后面 rank 的显存?
- 通信缓冲是不是按最坏情况留的?
一个把这四步走完的真实例子是 Kimi K3 的显存篇,那篇有一张可以拖并行度的账本。