大模型训练与推理的显存模型:参数、优化器状态与 KV Cache

很多 AI Infra 问题,最后都会回到一句非常朴素的话:显存里到底放了什么?如果这笔账没有算清,讨论 ZeRO、量化、PagedAttention 或多卡并行,很容易变成技术名词的堆砌。

本文不按“模型结构、训练优化、推理优化”的教材顺序展开,而是从一张统一的显存账本出发,把 Transformer 参数、优化器状态、激活值和 KV Cache 放进同一个分析框架。目标不是背住某个模型需要多少 GB,而是拿到任意配置后,都能判断它是否装得下、为什么装不下,以及应该用什么代价换取空间。

训练与推理显存账本

图 1:训练显存和推理显存应使用两套状态模型,但最终都要回到统一的资源预算。

一、先区分两个问题:装得下和跑得好

“一张 H100 能不能运行 7B 模型”不是一个完整问题。至少还要补充四个条件:

  • 训练还是推理;
  • 使用 BF16、FP16、FP8、INT8 还是 INT4;
  • Batch Size、序列长度和并发量是多少;
  • 是否使用张量并行、FSDP、量化或 CPU Offload。

模型勉强装进显存,也不代表系统可以稳定运行。CUDA Context、通信 Buffer、临时 Workspace、CUDA Graph 和框架缓存都需要空间。生产环境通常不会把显存利用率推到理论上的 100%,否则一次较长的请求或稍大的临时张量就可能触发 OOM。

因此容量规划要依次回答:

  1. 权重和长期状态占多少;
  2. 随输入变化的动态状态占多少;
  3. 运行时还需要留下多少安全空间;
  4. 达到目标吞吐时,动态状态会增长到什么程度。

1.1 用“生命周期 × 副本数”统一不同状态

仅按“权重、梯度、激活”分类还不够。更稳健的模型是同时标注状态的生命周期和复制方式:

状态 生命周期 主要增长变量 是否适合分片
模型权重 服务或训练进程全程 参数量、精度 TP/FSDP 可分片
优化器状态 整个训练任务 参数量、优化器 ZeRO/FSDP 可分片
激活 一个 Micro Batch 的前反向 Batch、Sequence、层数 SP/CP 或重计算
KV Cache 一个请求的存活时间 并发、已缓存 Token TP 分片、分页管理
Workspace 一次算子或一张计算图 Kernel、Shape、后端 通常不可简单均分
通信 Buffer Collective 执行期间 Bucket、并行组 与通信实现相关

对于任意长期状态,可以用下面的抽象检查是否重复计算:

Mstate=元素数×每元素字节数×副本因子分片因子M_{state}=\text{元素数}\times\text{每元素字节数}\times \frac{\text{副本因子}}{\text{分片因子}}

这个公式比“参数量除以 GPU 数”多了一个关键变量:副本因子。某个张量即使理论上可分片,也可能因为计算或通信需要,在特定阶段临时恢复完整副本。FSDP 的峰值 AllGather、TP 中未切分的小参数以及 CUDA Graph 私有池都属于这类情况。

二、从配置文件估算参数量

设隐藏维度为 dd,层数为 LL,词表大小为 VV,FFN 中间维度为 dffd_{ff}

2.1 Embedding 和 LM Head

输入 Embedding 的参数量为:

V×dV \times d

如果模型使用 Weight Tying,让输入 Embedding 和输出 LM Head 共享权重,这部分只存一份;否则还要再加一个 V×dV \times d

Decoder-only 模型由多个 Decoder Block 堆叠而成

图 2:参数账本必须覆盖 Embedding、每层 Attention/FFN 以及最终输出层,而不能只按隐藏维度粗略估计。

2.2 Attention

标准 Multi-Head Attention 包含 Q、K、V 和输出投影,忽略 Bias 时约为:

4d24d^2

GQA 或 MQA 会减少 K、V 投影的输出维度。设 Query Head 数为 HqH_q,KV Head 数为 HkvH_{kv},单头维度为 dh=d/Hqd_h=d/H_q,则每层 Attention 参数量约为:

d2+2d(Hkvdh)+d2d^2 + 2d(H_{kv}d_h) + d^2

GQA 对总参数量的影响通常没有对 KV Cache 那么显著,但会直接改变推理时每个 Token 的缓存大小。

2.3 SwiGLU FFN

当前很多 Decoder-only 模型使用 SwiGLU。它包含 Gate、Up 和 Down 三个矩阵,因此每层约为:

3d×dff3d \times d_{ff}

这也是为什么大模型中 FFN 往往比 Attention 占用更多参数。

2.4 一个可复用的近似式

忽略 Norm 和 Bias 等小项,一个 Decoder-only 模型可以近似写成:

NVd+L(Pattn+3ddff)N \approx Vd + L\left(P_{attn}+3dd_{ff}\right)

手算时不必追求个位数精度。容量规划更关心 7B、13B、70B 这种数量级,以及不同模块在总参数中的比例。

三、训练显存:为什么常见估算是每参数 16 字节

使用 BF16 参数和 AdamW 训练时,一种常见的静态显存账本如下:

项目 每参数字节数 说明
BF16 参数 2 B 前向与反向使用
BF16 梯度 2 B 实现不同可能有所变化
FP32 主权重 4 B 保证更新精度
Adam 一阶动量 4 B 梯度指数移动平均
Adam 二阶动量 4 B 梯度平方指数移动平均
合计 16 B 不含激活和临时 Buffer

因此,静态训练状态常用下面的数量级估算:

Mstatic16NM_{static} \approx 16N

一个 7B 模型仅静态状态就约为 112 GB。即使不考虑激活值,它也无法直接放进一张 80GB GPU 中进行标准 AdamW 训练。

需要特别注意:16N 是便于做决策的工程近似,不是所有框架都严格一致。梯度精度、优化器实现、参数 Flatten、通信 Bucket 和分配器行为都会影响真实结果。

3.1 激活值为什么更难估算

参数和优化器状态主要由模型规模决定,激活值则还受到以下因素影响:

  • Micro Batch Size;
  • 序列长度;
  • 层数和隐藏维度;
  • Attention 是否保存 S×SS\times S 中间矩阵;
  • 是否启用 FlashAttention;
  • 是否启用 Activation Checkpointing;
  • Tensor Parallel 和 Sequence Parallel 的切分方式。

激活显存通常可以抽象为:

MactB×S×d×LM_{act} \propto B \times S \times d \times L

不同实现的常数项差异很大,所以最可靠的方法是先做理论上界,再用框架的显存快照和 Profiler 验证。

四、推理显存:权重之外,真正会增长的是 KV Cache

推理时不再保存梯度和优化器状态,显存的主体变为:

Minference=Mweights+MKV+MruntimeM_{inference}=M_{weights}+M_{KV}+M_{runtime}

4.1 权重显存

权重部分可以直接近似为:

精度 每参数字节数 70B 权重量级
FP32 4 B 280 GB
BF16/FP16 2 B 140 GB
INT8 1 B 70 GB,加上量化元数据
INT4 0.5 B 35 GB,加上量化元数据

这张表解释了为什么 70B BF16 模型至少需要两张 80GB GPU,而 INT4 版本可能装入单张大显存 GPU。但“装入”不等于吞吐一定更高:量化 Kernel、反量化开销和硬件支持同样重要。

4.2 KV Cache 通用公式

每层都要为历史 Token 保存 K 和 V。设层数为 LL,KV Head 数为 HkvH_{kv},单头维度为 dhd_h,缓存精度字节数为 bb,则每个 Token 的 KV Cache 为:

MKV/token=2LHkvdhbM_{KV/token}=2LH_{kv}d_hb

总 KV Cache 为:

MKV=B×S×2LHkvdhbM_{KV}=B\times S\times 2LH_{kv}d_hb

其中 BB 是并发序列数,SS 是已经缓存的平均序列长度。这里最容易被忽略的是:KV Cache 不只随上下文长度增长,也随并发数线性增长。

以 32 层、32 个 KV Head、Head Dimension 为 128 的 BF16 模型为例:

2×32×32×128×2=524288 Bytes2\times32\times32\times128\times2=524288\text{ Bytes}

即每个 Token、每个请求约 0.5 MiB。当并发为 16、序列长度为 4096 时,KV Cache 约为 32 GiB。这还没有计入模型权重和运行时空间。

4.3 GQA 为什么对推理系统格外重要

如果 Query Head 为 64,而 KV Head 只有 8,那么相对于同样使用 64 个 KV Head 的 MHA,KV Cache 可以缩小到原来的八分之一。

这说明一个架构改动可能同时影响三个层面:

  • 模型层:多个 Query Head 共享 K/V;
  • Kernel 层:Attention 的读取和线程映射发生变化;
  • 系统层:同一张 GPU 可以容纳更多并发请求。

五、把显存优化理解为交换

几乎所有显存优化都不是免费的。

技术 主要节省对象 付出的代价
混合精度 权重、梯度、激活 数值稳定性管理
INT8/INT4 量化 推理权重与带宽 精度风险、专用 Kernel
Activation Checkpointing 训练激活 反向时重计算
ZeRO/FSDP 参数、梯度、优化器状态 更频繁的集合通信
CPU/NVMe Offload GPU 静态状态 PCIe/存储传输延迟
GQA/MQA KV Cache 模型表达能力与架构约束
PagedAttention KV 碎片和预留浪费 块表间接寻址与管理成本
Prefix Cache 重复前缀的 KV 与 Prefill 哈希、淘汰和命中率管理

做技术选型时,应该把问题改写成:“当前最稀缺的是显存、带宽、计算还是延迟?我愿意用什么资源交换?”

六、训练与推理的容量规划顺序

6.1 训练

  1. 根据配置估算参数量;
  2. 按优化器和精度计算静态状态;
  3. 估算激活上界;
  4. 判断是否必须进行参数分片;
  5. 根据机器拓扑安排 TP、FSDP/DP 和 PP;
  6. 留出通信 Buffer 与分配器碎片空间;
  7. 用一个较小 Batch 实测峰值显存,再反推可用配置。

6.2 推理

  1. 计算权重显存;
  2. 计算目标并发和上下文长度下的 KV Cache;
  3. 为 CUDA Graph、Workspace 和框架预留空间;
  4. 决定 TP 数量与量化格式;
  5. 决定 max_model_len、并发上限和显存利用率;
  6. 用真实 Prompt 长度分布压测,而不是只测固定短输入。

6.3 把公式固化成容量估算器

理论公式适合白板推导,工程上更适合把假设写入一个可以审计的小工具。下面的函数故意不调用任何深度学习框架,便于检查单位和输入:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
from dataclasses import dataclass

GIB = 1024 ** 3

@dataclass
class ModelConfig:
layers: int
hidden_size: int
query_heads: int
kv_heads: int
head_dim: int
parameters: int

def weight_gib(parameters: int, bytes_per_weight: float) -> float:
return parameters * bytes_per_weight / GIB

def training_static_gib(parameters: int, bytes_per_parameter: float = 16) -> float:
"""BF16 参数/梯度 + FP32 主权重 + Adam m/v 的常用近似。"""
return parameters * bytes_per_parameter / GIB

def kv_cache_gib(
cfg: ModelConfig,
batch_size: int,
cached_tokens_per_request: int,
bytes_per_element: int = 2,
) -> float:
elements = (
2
* cfg.layers
* cfg.kv_heads
* cfg.head_dim
* batch_size
* cached_tokens_per_request
)
return elements * bytes_per_element / GIB

这个估算器还没有考虑 TP 切分、量化 Scale、对齐、块内碎片和运行时 Workspace。它的意义不是取代框架,而是把“我们假设了什么”显式保存下来。实测峰值与理论值不一致时,就可以逐项寻找缺失状态。

七、技术 Q&A

Q1:为什么 7B 模型的 BF16 权重约 14GB,标准 AdamW 训练却可能超过 100GB?

14GB 只计算了 7B × 2 Bytes 的模型权重。混合精度 AdamW 训练还要保存 BF16 梯度、FP32 主权重、FP32 一阶动量和 FP32 二阶动量,常用静态近似为 16 Bytes/参数,即约 112GB。实际训练还需要激活、临时张量、通信 Bucket 和 CUDA Runtime,所以 112GB 仍不是完整峰值。

Q2:为什么 GQA 对参数量影响有限,却能显著降低 KV Cache?

GQA 只缩小 K、V 两个投影以及对应缓存,Attention 中 Q 和输出投影、FFN、Embedding 都没有同比缩小,所以总参数量下降有限。KV Cache 只保存每层 K/V,它与 kv_heads 成正比;KV Head 从 64 减到 8 时,这部分缓存理论上直接缩小八倍。

Q3:ZeRO-3、Activation Checkpointing 和 INT4 量化分别解决什么状态?

ZeRO-3 分片训练期的参数、梯度和优化器状态;Activation Checkpointing 减少训练期保存的激活,用反向重计算交换显存;INT4 量化主要压缩推理权重,并降低权重读取带宽。三者作用对象不同,不能互相替代,但可以组合。

Q4:理论容量估算应该如何与运行时测量闭环?

先用公式得到静态下界和动态上界,再在目标硬件上分别测启动后、Prefill 峰值、Decode 稳态和压力场景峰值。PyTorch 可结合 memory_allocatedmemory_reserved、Memory Snapshot 与 Profiler;推理引擎还应记录 KV Block 数、抢占次数和缓存命中率。理论模型负责解释,运行测量负责校准。

八、系列导航

下一篇:LLM 自回归推理的执行路径

九、参考资料

原始论文与技术报告

文档与进一步阅读