KV Cache 原理与显存优化:大模型推理 OOM 排查实战
发布时间:2026/9/18 11:47:02 作者:尧图编辑部 阅读量:1,286

前几天帮一个朋友看他本地部署的推理服务8G 显存的卡模型权重放进去还剩一点余量单条对话跑得好好的结果他把并发调到 8还没跑到第二轮就开始报显存不足。他把模型换小了一档问题照旧把max_tokens砍到 512还是撑不住。最后我让他把上下文长度从 32K 降到 4K瞬间就稳了。吃掉显存的不是权重是 KV Cache——那个在你眼皮底下悄悄长大的缓存。这篇东西我想把 KV Cache 从头到尾讲一遍它到底缓存了什么、为什么只缓存 K 和 V、token 一个一个往外蹦的时候它在干什么、显存账本怎么算、GPU 在这件事上到底卡在算力还是带宽。写给谁看写过一遍推理接口、被 OOM 折腾过、想搞清楚max_model_len和gpu_memory_utilization该怎么设的人。纯调 API 的也能看至少以后不会再对着显存曲线发懵。1. KV Cache 到底在解决什么问题1.1 大模型吐字为什么是一个一个来的自回归生成这件事本质上是把预测下一个词这个动作重复执行。你给它一段 prompt它输出第 1 个 token然后把刚生成的这个 token 拼回输入再去算第 2 个如此往复直到遇到结束符或者撞上长度上限。这个流程决定了输出的串行性第 n 个 token 必须等第 n-1 个算完才能开始。这里有个很多人一开始会误解的点prefill 和 decode 是两个性质完全不同的阶段。prefill 阶段处理你输入的整段 prompt几百上千个 token 一次性喂进去做的是大矩阵乘法GPU 的算力单元跑得满满当当这个阶段叫compute-bound。decode 阶段则完全相反每一步只处理一个新 token矩阵乘法退化成了矩阵-向量乘法GEMV算力单元大部分时间在发呆真正忙的是显存控制器这个阶段叫memory-bound。这两个阶段的差异直接决定了 KV Cache 的存在意义和它的代价。prefill 阶段我们关心的是算力和首 token 延迟TTFTdecode 阶段我们关心的是显存带宽和单 token 输出速度TPOT。KV Cache 是为 decode 阶段服务的也是 decode 阶段最大的显存开销来源。1.2 不做缓存会发生什么假设你在生成第 100 个 token。要算这个位置的注意力输出需要拿当前位置的 Query 去和前面 99 个位置的 Key 做点积得到权重后再对前面 99 个位置的 Value 做加权求和。问题来了前面 99 个位置的 K 和 V 是从哪来的它们是由那 99 个位置的输入向量分别乘上各自的权重矩阵算出来的。如果没有任何缓存每一次生成新 token你都得把前面所有 token 的 K 和 V 重新算一遍。第 1 步算 1 个第 2 步算 2 个第 3 步算 3 个……到第 n 步就要算 n 个。整个生成过程里K/V 投影的总计算量是 O(n²)而注意力矩阵的计算量累积起来是 O(n³)——因为第 t 步的注意力矩阵是 t×t 大小所有步骤加起来就是三次方级别。具体点说生成 2048 个 token不做缓存的话你要重复计算大约 200 万次 K/V 投影而实际上只有 2048 个是新的。这个浪费是灾难性的而且随着长度增长浪费比例越来越大长度翻倍重复计算量翻四倍。所以 KV Cache 不是什么锦上添花的优化它是自回归生成能不能实用的前提。1.3 加上缓存之后复杂度降到哪一档有了 KV Cache每一步只需要给新来的那 1 个 token 算它的 K 和 V然后追加到缓存里。K/V 投影的总计算量从 O(n²) 降到 O(n)。注意力部分每步的 Query 只有 1 个batch 内它要和长度为 t 的 K 缓存做点积所以每步 O(t)累积 O(n²)。整体从 O(n³) 降到 O(n²)。这个降幅在短序列上感觉不明显长序列上是质变。我做过一个粗略的对比测试同一个 7B 模型生成 4096 个 token不开缓存的话在 4090 上要跑十几分钟量级开了缓存是几十秒量级——差距不是几倍是两个数量级的体感差别这个数字受具体实现影响只是个数量级参考。不过缓存不是免费的。你省下了计算代价是把显存当成了计算结果的暂存区。省下来的算力换成了持续增长的显存占用。这就是后面所有显存问题的源头。2. 从 Attention 公式拆解 KV Cache 缓存的具体对象2.1 Q、K、V 三个矩阵各自在扮演什么角色标准的缩放点积注意力可以写成一行公式Attention(Q, K, V) softmax(Q K^T / sqrt(d_head)) V用信息检索的类比来理解最直观。你手里有一个查询词Query要去一个资料库里找匹配。资料库里的每份资料有两个属性一个是索引标签Key用来跟你的查询词算相似度一个是正文内容Value一旦相似度算出来就按相似度加权把这些正文混合起来。Query 和 Key 做点积得到的是这个查询跟每个位置有多相关的分数矩阵经过 softmax 归一化后变成权重再乘上 Value 得到最终输出。三者都是输入序列经过不同的线性投影W_q、W_k、W_v得到的形状都是[batch, seq_len, num_heads, head_dim]。关键在于这三个角色的时间特性完全不同。Query 描述的是当前位置想要什么它只跟当前这个位置有关。Key 和 Value 描述的是某个位置是什么它是位置的固有属性只要这个位置的内容不变、模型权重不变它算出来就不变。2.2 为什么只缓存 K 和 V不缓存 Q这是理解 KV Cache 最核心的一问我见过不少人卡在这。生成第 t 个 token 的时候你需要的 Query 只有 1 个就是位置 t 的那个。而你需要参与点积的 Key 有 t 个位置 1 到 t需要加权求和的 Value 也有 t 个。位置 t 的 K 和 V 是刚算的位置 1 到 t-1 的 K 和 V 是之前每一步算过、以后每一步还要继续用的。那 Query 呢位置 t-1 的 Query 在第 t-1 步用完之后第 t 步还需不需要它不需要。第 t 步的注意力是位置 t 的 Query去查位置 1 到 t 的 Key旧 Query 已经完成了它的历史使命再也不会被任何后续步骤引用。所以缓存 Q 是纯粹的浪费——存了也没人会读。一句话总结Q 是一次性消费品K 和 V 是长期资产。缓存只存长期资产。这里还有一个容易忽略的细节在 prefill 阶段输入的整段 prompt 里每个位置的 Q、K、V 都是要算的因为 causal mask 之下位置 i 的 Query 要和位置 1..i 的 Key 做注意力。但只有 K 和 V 会被留下来给后面的 decode 用Q 用完即弃。所以你会看到一个现象prefill 结束的那一刻KV Cache 的占用就已经等于prompt 长度那么多而不是从 0 开始慢慢涨。这一点在排查显存曲线时会很有用。2.3 缓存的张量长什么样、怎么增长从实现角度看KV Cache 通常是两块张量一块存 Key一块存 Value也有实现把两者拼在一起成一个张量减少一次内存分配。逻辑形状是[num_layers, batch, num_kv_heads, max_seq_len, head_dim]注意这里的max_seq_len是预留长度不是实际长度。绝大多数框架为了避开反复申请释放内存的开销会在请求开始时就按最大可能长度预分配一块连续显存然后用一个当前长度的计数器来标记有效区域。这就是为什么有时候你看到显存占用在请求一开始就跳到一个很高的值——那不是泄漏是预留。每生成一个 token做的事情就是把新的 K、V 写入到current_len这个位置然后current_len 1。写入的位置是连续递增的所以单条请求内部的追加操作非常轻量。注意事项预分配策略虽然避免了碎片但它意味着单条请求的实际显存占用往往按max_model_len算而不是按实际生成长度算。如果你把max_model_len设成 128K 但实际只用 2K等于白白浪费了 98% 的预留空间。这是很多人配错参数字后并发上不去的隐藏原因。3. 显存账本KV Cache 到底吃掉多少显存3.1 单个 token 的缓存字节数怎么算这个公式值得记下来我建议直接抄进笔记每 token 缓存字节数 2 × num_layers × num_kv_heads × head_dim × dtype_bytes拆开看每一项2Key 和 Value 各一份。num_layers每一层都有自己独立的 KV Cache层与层之间不共享。num_kv_heads注意是KV 头数不是注意力头数。用了 GQA 的模型这两者不一样差得还挺多。head_dim每个头的维度通常是hidden_size / num_headsLlama 系是 128。dtype_bytesfp16/bf16 是 2fp8 是 1int4 压到 0.5 左右含 scale 开销会略高。如果要算一条完整序列的总占用再乘上序列长度和 batch 大小总缓存字节 2 × L × H_kv × d_head × S × B × dtype_bytes其中S是序列长度prompt 生成B是并发请求数。这两个乘数就是显存爆炸的放大器。3.2 拿几个常见模型手算一遍理论讲完不算数算一遍才有感觉。下面这些都是按 fp16 算的截图数据是我自己按公式推的跟实测会有几个百分点的偏差主要来自框架的预留和对齐。模型层数KV 头数head_dim每 token 缓存8K 上下文单条Llama-2-7B (MHA)3232128512 KB4.0 GBLlama-3-8B (GQA)328128128 KB1.0 GBQwen2.5-7B (GQA)28412856 KB0.44 GBQwen2.5-32B (GQA)648128256 KB2.0 GBLlama-2-7B 那一行很能说明问题每 token 512 KB一条 8K 的对话就要 4 GB权重本身才 13 GB 左右fp16。如果并发 4 条KV Cache 直接吃掉 16 GB比权重还多。这就是老一代 MHA 模型在长上下文场景下特别吃显存的原因。再看 Qwen2.5-7B56 KB/token同样 8K 上下文只要 0.44 GB。差了将近 10 倍原因就是 KV 头数从 32 降到 4。模型结构上的一个改动直接把显存账本改写了。再算个极端点的32B 模型跑 32K 上下文256 KB × 32768 8 GB单条请求。这时候权重是 64 GB加 8 GB 缓存再加激活值和框架开销80G 卡上并发基本只能到 1-2。3.3 为什么缓存比权重更容易成为爆点权重的占用是固定且可预测的加载完就定死了你知道它占多少。KV Cache 是动态且乘性增长的它随并发线性增长随上下文长度线性增长两个维度一乘增长是双倍的。而且 KV Cache 有一个讨厌的特性——它不能像权重那样被高效换出。权重是只读的多个请求共享同一份KV Cache 是每个请求私有的读写频繁换到主机内存再换回来PCIe 的带宽和延迟会直接把 decode 速度拖垮。所以实际的显存规划里权重只是下限KV Cache 才是决定你能塞多少并发的那个变量。我一般的粗略分配是这样的先给权重留出参数量 × dtype_bytes再给 CUDA 上下文和框架开销留 1-2 GB剩下的显存乘个 0.8-0.9 就是 KV Cache 的可用池子。然后用池子大小 / (每 token 字节 × 平均序列长度)估算能同时跑几条。这个估算在调gpu_memory_utilization的时候非常好用。4. GPU 侧的真实瓶颈算力还是带宽4.1 算术强度告诉你 decode 阶段卡在哪判断一个 kernel 卡在算力还是带宽看算术强度Arithmetic Intensity也就是每搬运 1 字节数据能换多少次浮点运算单位是 FLOP/Byte。把它跟 GPU 的拐点比一下就有结论。以 RTX 4090 为例fp16 密集算力约 82 TFLOPS显存带宽约 1008 GB/s拐点大约在82e12 / 1008e9 ≈ 81 FLOP/Byte。也就是说算术强度低于 81 的负载都是 memory-bound。现在算 decode 阶段生成一个 token 的实际情况。以 8B 模型 fp16 为例权重要完整读一遍 16 GB假设 KV Cache 有 2 GB 也要读一遍总共约 18 GB 的数据搬运。计算量大约是2 × 8e9 16 GFLOP每个参数两次运算。算术强度 16e9 / 18e9 ≈ 0.89 FLOP/Byte。跟拐点 81 差了将近90 倍。这意味着在 decode 阶段GPU 的算力单元有超过 98% 的时间在空转真正的时间全部花在把数据从显存搬到计算单元上。用时间验证一下搬运 18 GB 需要18 / 1008 ≈ 17.9 ms计算 16 GFLOP 需要16e9 / 82e12 ≈ 0.2 ms。搬运是计算的 90 倍。这 17.9 ms 就对应理论上限约 56 token/s 的单条输出速度实际还要打折。所以 decode 阶段的优化方向从来不是算得更快而是少搬点数据或搬得更有效率。4.2 带宽、碎片和 batch 之间的三角关系既然瓶颈是带宽那提高吞吐的路子就是用更多的计算来摊薄搬运成本。具体做法就是加大 batch多条请求共享同一份权重读取权重读一次能服务 B 条请求。这时候算术强度乘以 B当 B 足够大时负载就从 memory-bound 往 compute-bound 迁移。但 batch 一开大KV Cache 的占用又跟着涨。这就形成了一个闭环约束batch 受显存限制显存被 KV Cache 占用KV Cache 又正比于 batch 和序列长度。你能跑多大 batch取决于你能省下多少 KV Cache 显存。除了容量还有两个隐形的敌人。一个是显存碎片早期的实现给每条请求按最大长度预分配连续空间实际用了 30% 的话剩下 70% 就空着谁也用不了。另一个是不规则长度一批请求里有的已经生成了 3000 token有的才 50 token按最长的对齐就会浪费大量空间。这两点加起来在真实负载下显存的有效利用率可能只有 20%-40%。4.3 分页管理是怎么把浪费压下去的vLLM 提出的 PagedAttention 是这方面最出名的一个思路核心灵感来自操作系统的虚拟内存分页。做法是把 KV Cache 切成固定大小的block常见是 16 个 token 一块block 之间不要求物理连续。每个请求维护一张 block table记录我的第几个逻辑块存在哪个物理块里。注意力计算时kernel 按 block table 去非连续的位置取 K 和 V。这样一来显存分配从按最大长度预留变成了按需一页一页地领。请求生成到哪就领到哪最后一个块没用满也没关系浪费的上限就是块内未使用的部分平均一半的块大小。同时块的物理位置灵活碎片问题自然消失。公开的数据里这种方案能把显存浪费从 60%-80% 压到 4% 以下同一张卡上能跑的并发数因此翻好几倍。另一个附带的好处是前缀共享。如果多条请求的 system prompt 完全一样它们的 K 和 V 在前缀部分算出来是完全相同的那就可以让这些请求的 block table 指向同一批物理块只在分叉之后各自独立。多轮对话场景下这个收益非常大因为历史对话在下一轮里就是共同前缀。我实测过一个带长 system prompt 的场景开前缀缓存后首 token 延迟下降了大概 40%显存占用也明显下来了。5. 工程上的压缩与优化手段5.1 从模型结构下手MQA 和 GQA最省事的优化是在训练阶段就做掉的也就是减少 KV 头数。MHA多头注意力里Q、K、V 的头数一样多KV Cache 最大。MQA多查询注意力把所有的 Q 头共享同一组 K 和 VKV 头数降到 1缓存直接砍到 1/NN 是头数代价是表达能力有损失质量会掉一些。GQA分组查询注意力是折中把 Q 头分成 G 组每组共享一组 K/VKV 头数从 N 降到 G。Llama-3-8B 是 32 个 Q 头、8 个 KV 组比例 4:1缓存省了 75%。Qwen2.5-7B 更激进28 个 Q 头配 4 个 KV 组比例 7:1。这就是为什么新一代模型在长上下文上比老模型从容得多。不过这个选择在你拿到模型权重的那一刻就固定了推理侧改不了。你能做的只是在选模型时有意识地看这个指标。如果你要部署的场景是长上下文加高并发num_kv_heads的重要性不比参数量低。有些模型在配置里写的num_key_value_heads参数就是它。5.2 KV Cache 量化收益、代价和踩坑点如果模型结构动不了那就对缓存本身做量化。fp16 是 2 字节量化到 int8 就是 1 字节缓存减半量化到 int4 理论上减到 0.5 字节但通常要额外存 scale 和 zero_point实际大约 0.55-0.65 字节每元素。收益很直接原来只能跑 4 并发的现在能跑 7-8。代价是精度损失而且 KV 量化的精度损失跟权重量化不一样它影响的不是单个权重而是注意力权重的分布。我踩过的一个坑是短序列上测不出问题一上长上下文就开始胡言乱语。原因是 K 的数值分布会随着序列变长出现明显的离群通道某些维度上的值特别大per-tensor 量化会被这些离群值带偏把正常值全部压到很窄的量化区间里。解决办法是用 per-channel 或者 per-token 的量化粒度或者干脆只对 V 做量化、K 保持 fp16因为 V 的分布通常温和得多。实操建议KV 量化不要一上来就上 int4。先在 int8 上跑一遍你的实际业务 prompt对比输出质量做个 A/B。int8 通常质量损失在可接受范围int4 就要谨慎了尤其是需要精确数值推理、代码生成这类任务。5.3 前缀缓存、滑动窗口与驱逐策略除了压缩还有少存这条路。滑动窗口注意力是让每个位置只关注最近 W 个 token。缓存就只需要保留最近 W 个位置的 K 和 V超出的直接扔掉显存占用从O(S)变成O(W)是个常数。Mistral 早期版本用的就是这个代价是对超长距离依赖的建模能力变弱。驱逐策略是另一种思路缓存满了按某种规则淘汰一部分。常见的有 H2O 这类基于注意力分数的重要性评估——注意力权重长期很低的位置说明它对后续输出影响小可以优先扔掉。还有些实现会保留最开头的几个 token所谓的 attention sink因为实测发现丢掉序列开头的 token 会导致质量明显崩塌哪怕它们的注意力分数不高。前缀缓存前面提过多请求共享相同前缀的 KV。这个在 API 服务里几乎是必开项尤其是 system prompt 很长或者多轮对话的场景。这几种手段可以叠加。我一般的选择顺序是先确认模型本身用了 GQA再开前缀缓存零质量损失然后开分页管理零质量损失最后才考虑量化和驱逐有质量损失需要评估。6. 实操怎么测出真实的显存占用6.1 观测工具与观测点先解决看得见的问题。nvidia-smi是最基础的一层但它给的是进程级的显存占用看不到细分。想看细分得用 PyTorch 自己的接口import torch # 当前实际被张量占用的显存 print(torch.cuda.memory_allocated() / 1024**3, GB) # 被缓存分配器持有、但当前没被张量使用的显存 print(torch.cuda.memory_reserved() / 1024**3, GB) # 完整的分配快照能看到各段占用 print(torch.cuda.memory_summary())这里有个关键区别要搞清楚memory_allocated是真正被张量用掉的memory_reserved是 PyTorch 缓存分配器向驱动申请下来、但可能暂时闲置的。OOM 往往发生在 reserved 已经很大、但 allocated 不大的时候——不是没内存是内存被缓存分配器攥着不放。想释放可以调torch.cuda.empty_cache()但只对 reserved 部分有效且会带来性能抖动生产环境慎用。实际部署里我更习惯在服务启动后先打一次基线快照加载完权重、还没接请求时的显存占用这个数就是底噪。之后每接一批请求再打一次两者一减就是 KV Cache 的实际开销。这个差值比任何理论公式都准。6.2 参数调节的顺序和经验值显存不够的时候参数不能乱调。我摸索出来的顺序是这样的优先级参数调整方向影响1max_model_len降到业务实际需要显存占用线性下降无质量损失2gpu_memory_utilization提到 0.90-0.95提高可用池子但要留余量3max_num_seqs适当降低直接限制并发牺牲吞吐4KV 量化int8 起步缓存减半需评估质量5换更小的模型最后手段质量下降明显第一条为什么排最前因为max_model_len是乘法因子把它从 32K 降到 8KKV Cache 直接少 75%而且一点质量都不损失只要你的业务确实用不到那么长。我见过太多人为了以防万一设了个超长的上下文上限结果把并发能力全吃掉了。gpu_memory_utilization这个参数在 vLLM 里指的是允许使用的显存比例。设 0.9 意味着框架会在启动时就把 90% 的显存划走作为 KV Cache 池子。设太高比如 0.98有风险因为推理过程中还有激活值、临时张量、CUDA kernel workspace 要分配容易在某个瞬间撞上 OOM。我的经验是 0.85-0.92 之间比较安全具体看模型和序列长度。6.3 常见问题速查现象可能原因排查动作启动就 OOM权重 预分配 KV 池超出显存降gpu_memory_utilization或max_model_len跑一阵才 OOM长序列请求累积并发超出池子看请求长度分布降max_num_seqsreserved 高但 allocated 低缓存分配器碎片检查是否有频繁变长张量分配单条很快并发一上就慢batch 增大导致带宽争抢对比不同 batch 下的 TPOT显存够但吞吐上不去decode 阶段 memory-bound加大 batch 摊薄权重读取长上下文输出质量突降KV 量化在长序列下失效关量化对比或改 per-channel这张表里的每一行我基本都实际遇到过。最典型的是第二行跑一阵才 OOM启动正常压测正常上线跑了半小时开始报错。原因是短请求快速释放长请求持续占用池子被慢慢蚕食。解决办法不是加显存而是给请求长度设上限或者对超长请求做排队。7. 我在实际部署里踩过的几个坑说几个文档里一般不写的东西。第一个坑是 prefill 阶段的显存尖峰。我一直以为显存是随 decode 平滑增长的直到有一次看到显存曲线在请求刚进来的瞬间跳了一下。原因是 prefill 处理长 prompt 时中间激活值尤其是注意力分数矩阵会临时占用一大块显存长度是S²级别的。一个 8K 的 prompt注意力矩阵如果全物化尺寸相当可观。所以即使你的 KV Cache 池子留够了也可能在 prefill 那一刻被顶爆。解决办法是开 chunked prefill把长 prompt 切成小段处理牺牲一点 TTFT 换平稳。第二个坑是max_tokens和实际占用不匹配。有次我把max_tokens设为 4096以为显存是按实际生成长度算的。实际上框架按max_model_len预留max_tokens只是限制生成上限不影响预留。所以你把max_tokens调小对显存几乎没帮助——真正有用的是max_model_len。第三个坑是并发测试的方法不对。用固定长度的请求压测得出的并发上限偏乐观。真实负载里请求长度是长尾分布几个超长请求就能把池子占满其他请求只能排队。做容量规划时我会按 P95 甚至 P99 的长度去算而不是平均值。经验公式是可用并发 ≈ 池子大小 / (每 token 字节 × P95 长度)再打个 0.8 的折。第四个坑是忽略了 CUDA 上下文的固定开销。每个进程初始化 CUDA 上下文要占几百 MB多进程部署时这个开销会累积。还有些框架会在启动时做一次显存探测把整卡显存都摸一遍。所以理论能装下和实际能跑起来之间永远要留 1-2 GB 的余量。第五个坑有点反直觉显存利用率高不等于性能好。我曾经把池子塞到 95%结果 decode 速度反而下降了。原因是显存快满的时候分配器的行为会变差容易出现分配失败重试和碎片整理。留出 10% 左右的空闲系统整体更稳。这个跟操作系统磁盘别占满是一个道理。最后分享一个我常用的快速估算手法。拿到一个新模型先算单 token 缓存字节数公式前面给过。然后看你的卡剩多少显存给缓存除以单 token 字节得到能缓存多少 token。再用这个数除以你的目标上下文长度就是大概能同时跑几条。这个心算过程只要三十秒但能帮你在配置之前就判断出这个模型这张卡到底能不能干这个活省下很多来回试错的时间。