先记住这个答案
KV cache显存计算公式为:每token KV字节数 = 2 × 层数 × KV头数 × 头维度 × 精度字节数。总显存占用 = 每token字节数 × 序列长度 × 并发数。决定最大并发时,先用单卡显存减去模型权重和激活预留,再除以单请求的KV占用(每token字节数 × 最大上下文长度)。实际部署需考虑动态长度和碎片,通常预留20%余量。
- 每token KV显存 = 2×层数×KV头数×头维度×字节数
- 单请求KV占用 = 每tokenKV × 最大序列长度
- 最大并发 = (可用显存 - 权重 - 激活) / 单请求KV
KV cache显存公式推导
Transformer自回归解码时,每个token在每一层都要产生Key和Value向量,它们被后续token反复查询,因此必须缓存。设模型有L层,每层有H个KV头(MHA时等于注意力头数),每个头的维度为D,则每个token每层的KV向量元素数为2×H×D。乘以精度字节数B(FP16为2),得到每token每层字节数,再乘L得每token总字节数。
总KV显存 = 每token字节数 × 序列长度 × 并发数。序列长度指输入和输出总长度。在容量规划中,通常用最大上下文长度作为每个请求的保守估计,这样最大并发 = 可用显存 / (每token字节数 × 最大序列长度)。若使用GQA或MQA,KH头数会小于注意力头数,公式中H应替换为KV头数,显存显著下降。
7B模型并发估算实例
设模型为7B,L=32层,注意力头数=32,head_dim=128,所以隐藏维度=4096,KV头数=32(MHA)。FP16精度B=2。每token KV字节数 = 2×32×32×128×2 = 524288字节,即512KB。上下文窗口2K token,则单请求最大KV占用=512KB×2048=1GB。
在一张80GB的A100上,模型权重FP16约14GB,激活和临时缓冲预留6GB,那么可用KV显存=80-14-6=60GB。最大并发≈60GB/1GB=60。但实际预填和生成阶段共用KV,且不同请求长度不一,直接按60会引入风险。若仅将上下文降至1K(保持MHA),单请求KV变为512MB,并发约为120;若保持2K上下文但改用GQA(KV头数8),单请求KV降至256MB,并发可到240。
公式失效与修正
该公式假设所有请求都达到最大序列长度,且批处理中每个请求的KV独立分配。当使用PagedAttention时,显存按页管理,内部碎片减少但页表本身有开销,且物理块大小限制分配粒度。若请求平均长度远小于最大长度,按最大长度计算会严重低估并发,此时应基于P99或P95长度估算。
分块预填(chunked prefill)会同时处理多个请求的prefill,导致KV写入更动态,容量规划需在时间维度上考虑峰值。此外,模型并行(张量并行)时每卡只存一部分KV,公式需除以张量并行数。实际部署建议用profiling工具测量峰值KV占用,再反推算力。
容易答错的地方
- 混淆参数量与层数影响
- 很多人以为KV cache大小和模型参数量成正比,实际上它只取决于层数、KV头数和头维度。参数量由隐藏维度和层数决定,但MHA下头数×头维度等于隐藏维度,所以两者有相关性,但GQA会打破该关系。
- 忽略激活和碎片余量
- 直接拿满显存除以单请求KV得到并发,没有预留激活和显存碎片,导致运行OOM。正确做法是先从总显存中扣除模型权重、激活和安全余量,再计算可用KV空间。
面试官还会怎么问?
GQA头数为8时,KV cache能省多少?
如果原本MHA头数32,改为GQA8,则KV头数降为原来的1/4,每token KV字节数变为原来的1/4,同样上下文下单请求占用降至1/4,最大并发提升4倍。
如何测量线上真实的KV cache占用?
可通过NVIDIA DCGM或PyTorch profiler监控显存增长曲线,在固定负载下观察稳态占用。更直接的是从推理引擎(如vLLM)的日志中读取KV cache使用率。
为什么有时显存没用完却OOM?
显存碎片化导致没有连续大块分配。PagedAttention将KV分成固定大小页,减少碎片,但页表本身需显存。另外,prefill阶段临时激活也可能瞬间占大显存。
参考资料
示例用于理解所注明的运行环境与边界;延伸学习可结合原文中的更多案例。