1. 从一次推理延迟排查说起:KV Cache 到底是什么
前阵子帮朋友看一个文本生成的推理服务,现象很典型:单条请求响应挺快,一旦并发上来,延迟直接飙到没法用,GPU 显存还涨得厉害。把 profiling 打开一看,绝大部分时间耗在注意力计算上,而且每生成一个 token,前面所有 token 的键和值都被重新算了一遍。问题就出在这里——KV Cache(Key-Value Cache)没做好,或者说根本没意识到它在自回归生成里的分量。
先把结论摆在前面:KV Cache 是大模型自回归推理阶段的一种缓存机制。它把每一层注意力里已经算过的 Key 和 Value 张量存下来,生成新 token 时直接复用,不再重复计算历史部分。就这么一个动作,能把生成第 n 个 token 的注意力计算量从 O(n²) 降到 O(n),实际推理速度提升往往是几倍到十几倍。代价是显存占用随序列长度线性增长,这也是为什么长上下文场景下显存总是紧张。
这篇内容适合三类人看:一是刚接触 Transformer 推理、搞不清 prefill 和 decode 区别的入门者;二是正在做推理服务、被延迟和显存两头夹击的工程同学;三是想搞明白“为什么大模型吐字速度会越来越慢”的普通使用者。我会从注意力机制的基本计算讲起,把 KV Cache 为什么存在、怎么算、显存怎么估、坑在哪里,一层层拆开。涉及到的参数计算我会给出具体数字,代码部分用 PyTorch 风格示意,方便你直接对照自己的实现。
需要提前说明的是,下面关于工程实现的部分,比如分页管理、量化缓存这些,是基于当前主流推理框架的常见做法做的合理补充,不是某一家的私有方案,你可以按自己用的框架去对应。
2. 为什么需要 KV Cache:自回归生成的重复计算问题
2.1 自回归生成的基本流程
要理解 KV Cache,得先接受一个前提:现在主流的大语言模型都是自回归(autoregressive)生成的。也就是说,模型一次只吐一个 token,然后把新吐出来的 token 接到输入后面,再预测下一个,如此循环。你看到的一段几百字的回答,背后是几百次前向传播。
每一次前向传播,输入都是“原始 prompt + 已经生成的所有 token”。假设 prompt 有 100 个 token,已经生成了 50 个,那第 51 次前向传播的输入长度就是 150。注意,这 150 个 token 里,前 149 个在上一轮已经算过了,只有最后一个是新的。
问题来了:如果每一轮都把 150 个 token 完整地过一遍注意力,那前 149 个 token 的 Key 和 Value 就被反复计算了 51 次。这就是纯粹的浪费。序列越长,浪费越夸张。
2.2 注意力机制里到底算了什么
我们回顾一下缩放点积注意力。给定输入 X,通过三个线性投影得到 Query、Key、Value:
Q = X @ W_q K = X @ W_k V = X @ W_v然后注意力输出是:
Attention(Q, K, V) = softmax(Q @ K^T / sqrt(d_k)) @ V关键在于,对于第 t 个位置的 token,它要算注意力,需要的是所有位置(1 到 t)的 K 和 V,以及自己这个位置的 Q。它不需要历史位置的 Q,因为 Q 只用来和 K 做匹配,每个位置各算各的。
所以,历史 token 的 Q 算完就可以扔了,但历史 token 的 K 和 V 后面每一步都要用。这就是 KV Cache 名字的由来——只缓存 K 和 V。
2.3 不做缓存时的计算量
假设序列长度为 n,隐藏维度为 d,层数为 L。单层注意力里,Q、K、V 的投影各是 O(n·d²),注意力矩阵 Q@K^T 是 O(n²·d)。生成第 n 个 token 时,如果全量重算,这一层就要做 O(n²·d) 的矩阵乘法。
整个生成过程从 1 到 n,总计算量是 O(n³·d) 量级(每步 O(n²·d),共 n 步)。做了 KV Cache 之后,每步只需要算新 token 的 Q、K、V,以及新 Q 和缓存 K 的注意力,单步是 O(n·d),总共 O(n²·d)。差距是 n 倍。n 等于 1000 的时候,就是 1000 倍的差距,这还没算上内存带宽和 kernel 启动的开销。
提示:这里说的“计算量”是理论 FLOPs。实际推理里,decode 阶段往往是内存带宽受限而不是算力受限,KV Cache 减少的不只是计算,更重要的是减少了对历史 K、V 的重复读写。
2.4 一个直观的类比
你可以把自回归生成想象成滚雪球。每滚一圈,雪球变大一点。如果不做缓存,相当于每滚一圈都要把整个雪球重新捏一遍;做了缓存,相当于只把新粘上的那层雪捏上去,原来的球体保持不动。雪球越大,省下的力气越多。KV Cache 就是这个“保持不动的球体”。
3. KV Cache 的核心原理与显存账本
3.1 缓存的结构:每层、每头、每位置
KV Cache 不是一份,而是每一层都有一份。因为 Transformer 每一层的注意力参数不同,算出来的 K、V 也不同,不能跨层复用。
在每一层内部,如果是多头注意力(Multi-Head Attention),每个头也有自己独立的 K、V。所以缓存的形状大致是:
[num_layers, 2, batch_size, num_heads, seq_len, head_dim]其中那个 2 就是 K 和 V。有些实现会把 K 和 V 分开存成两个张量,有些会拼在一起,逻辑上一样。
3.2 显存占用怎么估
这是工程上最关心的问题。单个 token 的 KV Cache 大小可以这样算:
每 token 字节数 = 2 × num_layers × num_heads × head_dim × dtype_bytes注意num_heads × head_dim通常等于隐藏维度 d(在标准多头里)。所以也可以写成:
每 token 字节数 = 2 × num_layers × d × dtype_bytes举个具体例子。假设一个模型 L=32 层,隐藏维度 d=4096,用 FP16(2 字节)存储:
每 token = 2 × 32 × 4096 × 2 = 524288 字节 ≈ 0.5 MB也就是说,每生成一个 token,KV Cache 就要多占 0.5 MB。如果上下文长度到 8192,batch size 为 1,那就是:
0.5 MB × 8192 ≈ 4 GB这还只是 batch=1。如果并发 16 路,直接 64 GB 显存没了。这就是为什么长上下文 + 高并发是显存杀手。
| 参数 | 符号 | 示例值 |
|---|---|---|
| 层数 | L | 32 |
| 隐藏维度 | d | 4096 |
| 数据类型 | dtype | FP16 (2B) |
| 每 token 缓存 | - | 0.5 MB |
| 序列长度 | n | 8192 |
| batch size | b | 1 |
| 总缓存 | - | 约 4 GB |
3.3 prefill 和 decode 两个阶段的差异
KV Cache 的引入,把推理明确分成了两个阶段:
Prefill 阶段:处理用户输入的 prompt。这时候所有 token 都是已知的,可以并行计算,一次性把整段 prompt 的 K、V 算出来填进缓存。这个阶段是计算密集型,GPU 利用率高。
Decode 阶段:逐个生成新 token。每步只算一个新 token 的 Q、K、V,然后拿新 Q 去和缓存里所有 K 做注意力。这个阶段是内存带宽密集型,因为每步都要把整个 KV Cache 读一遍。
这两个阶段的性能特征完全不同,优化手段也不一样。很多推理框架会把它们分开调度,甚至用不同的 kernel。理解这一点,对排查“为什么首 token 慢”和“为什么后续吐字慢”很有帮助。
3.4 为什么只缓存 K 和 V,不缓存 Q
前面提过,Q 是“当前 token 用来查询的向量”,每个位置只在它自己被生成的那一步用一次,之后再也不用了。缓存 Q 没有任何复用价值,只会白白占显存。而 K 和 V 是“被查询的对象”,会被后续所有 token 反复读取,所以必须缓存。
这个设计不是随便定的,是从注意力计算的数学结构里推出来的必然结果。你只要记住一句话:Q 是一次性的,K 和 V 是长期被引用的。
4. 动手实现:从零写一个带 KV Cache 的注意力
4.1 不带缓存的朴素实现
先看一个最朴素的单步注意力,方便对比:
import torch import torch.nn.functional as F def attention_no_cache(x, W_q, W_k, W_v, past_kv=None): # x: [batch, seq_len, d] Q = x @ W_q K = x @ W_k V = x @ W_v d_k = Q.size(-1) scores = Q @ K.transpose(-2, -1) / (d_k ** 0.5) attn = F.softmax(scores, dim=-1) out = attn @ V return out每次调用都要把完整序列传进来,K、V 全量重算。生成 100 个 token 就调用 100 次,每次序列都在变长。
4.2 带 KV Cache 的实现
改造思路很简单:把历史 K、V 存起来,每次只算新 token 的 Q、K、V,然后把新 K、V 拼到缓存后面。
def attention_with_cache(x_new, W_q, W_k, W_v, past_kv=None): # x_new: [batch, 1, d] 只包含当前新 token Q = x_new @ W_q K_new = x_new @ W_k V_new = x_new @ W_v if past_kv is not None: K_past, V_past = past_kv K = torch.cat([K_past, K_new], dim=1) V = torch.cat([V_past, V_new], dim=1) else: K, V = K_new, V_new d_k = Q.size(-1) scores = Q @ K.transpose(-2, -1) / (d_k ** 0.5) attn = F.softmax(scores, dim=-1) out = attn @ V new_kv = (K, V) return out, new_kv关键点有三个:一是输入只传新 token,二是用torch.cat把新 K、V 接到历史后面,三是把更新后的缓存返回出去,供下一步使用。
4.3 多头版本要注意的维度
多头注意力的缓存形状是[batch, num_heads, seq_len, head_dim]。拼接的时候要拼在seq_len那一维,也就是dim=2,别拼错。很多新手第一次写会拼到 head 维度上,结果注意力全乱。
# K_past: [batch, num_heads, past_len, head_dim] # K_new: [batch, num_heads, 1, head_dim] K = torch.cat([K_past, K_new], dim=2)4.4 预分配缓存 vs 动态拼接
上面用torch.cat是最直观的写法,但工程上很少这么干。因为cat每次都会新分配一块内存并拷贝,序列长了之后开销很大,还会造成显存碎片。
主流做法是预分配一块固定大小的缓存,形状是[batch, num_heads, max_seq_len, head_dim],然后用一个位置指针记录当前写到哪了。每步只往对应位置写,不重新分配。
# 预分配 cache_k = torch.empty(batch, num_heads, max_len, head_dim, device=device) cache_v = torch.empty_like(cache_k) pos = 0 # 写入 cache_k[:, :, pos:pos+1, :] = K_new cache_v[:, :, pos:pos+1, :] = V_new pos += 1 # 读取时只取前 pos 个 K = cache_k[:, :, :pos, :] V = cache_v[:, :, :pos, :]这个改动看起来小,但对吞吐的影响很大。预分配避免了反复分配释放,也让内存访问更连续。
注意:预分配的最大长度要提前定好。如果实际序列超过这个长度,要么截断,要么扩容。扩容时机的选择是个工程权衡,扩太早浪费显存,扩太晚触发重分配卡顿。
4.5 位置编码的配合
KV Cache 缓存的是 K、V 的数值,但位置信息是在算 K、V 之前就注入进去的(比如 RoPE 旋转位置编码)。所以缓存里的 K 已经带了位置信息,后续直接复用没问题。但如果你用的是绝对位置编码,且实现方式是“在注意力分数上加位置偏置”,那就要小心:新 token 的 Q 和缓存 K 做分数计算时,位置偏置要按各自的绝对位置来算,不能简单复用。
RoPE 之所以在长上下文模型里流行,一个原因就是它和 KV Cache 配合得很自然——位置信息编码在 K 里,缓存即用,不需要额外处理。
5. 工程实践中的坑与优化手段
5.1 显存不够怎么办:几个方向
显存是 KV Cache 最直接的约束。常见应对手段有这么几类:
减少并发:最粗暴,但吞吐直接掉。适合延迟敏感、吞吐不敏感的场景。
缩短上下文:限制 max_seq_len,或者做滑动窗口注意力,只保留最近 N 个 token 的缓存。代价是丢失远距离信息。
量化缓存:把 KV Cache 从 FP16 降到 INT8 甚至 INT4。显存直接减半或减到四分之一,精度损失通常可控,但需要校准。这是目前性价比很高的手段。
分组查询注意力(GQA):让多个 Query 头共享一组 K、V 头。比如 32 个 Q 头只配 8 个 KV 头,缓存直接降到四分之一。现在很多开源模型默认就用 GQA,就是为了省这块显存。
分页管理:借鉴操作系统的虚拟内存思路,把缓存切成固定大小的块(page),按需分配,减少碎片。这个思路在主流推理框架里已经很成熟。
| 手段 | 显存收益 | 主要代价 |
|---|---|---|
| 减少并发 | 线性 | 吞吐下降 |
| 滑动窗口 | 与窗口成正比 | 丢失长距离信息 |
| 量化缓存 | 2x ~ 4x | 精度损失、需校准 |
| GQA | 与头数比成正比 | 表达能力略降 |
| 分页管理 | 减少碎片 | 实现复杂度 |
5.2 缓存复用:前缀共享
如果你的服务里有很多请求共享同一段前缀(比如相同的系统提示词),那这段前缀的 KV Cache 是可以复用的。第一个请求算完,把前缀部分的缓存留下来,后续请求直接接上,省掉重复的 prefill。
这个优化在多轮对话里特别有用:历史对话的缓存可以保留,新一轮只 prefill 新增的用户输入。实现上需要一个前缀匹配和缓存索引机制,复杂度不低,但收益很可观。
5.3 常见问题速查
| 现象 | 可能原因 | 排查方向 |
|---|---|---|
| 生成越来越慢 | 缓存没生效,每步全量重算 | 检查是否每步都传了完整序列 |
| 显存随长度暴涨 | 缓存未预分配,碎片严重 | 看显存分配曲线是否锯齿状 |
| 输出乱码/重复 | 缓存拼接维度错 | 检查 cat 的 dim 是否为 seq 维 |
| 首 token 慢后续快 | prefill 计算密集 | 正常现象,优化 prefill kernel |
| 并发上不去 | 缓存占用过大 | 估算每 token 字节数,考虑量化或 GQA |
| 长文本质量下降 | 滑动窗口截断了关键信息 | 调整窗口大小或换注意力方案 |
5.4 几个容易忽略的细节
第一,缓存的清理时机。请求结束后要及时释放对应的缓存块,否则显存会慢慢泄漏。分页管理里通常有个引用计数,归零就回收。
第二,batch 内不同序列长度不一致。同一个 batch 里,有的请求已经生成 100 个 token,有的才 10 个。这时候缓存的有效长度不同,注意力计算要做 mask,否则短序列会读到别的序列的缓存。这个 bug 很隐蔽,输出可能看起来“差不多对”,但质量会悄悄下降。
第三,数值精度。缓存长时间累积,如果中间有精度损失,误差会逐步放大。FP16 缓存在超长序列下可能出现数值不稳定,有些场景会保留一份 FP32 的累加。这个要看具体模型和任务,不是所有场景都需要。
第四,beam search 下的缓存管理。beam search 会同时维护多条候选路径,每条路径有自己的缓存。beam 合并、剪枝的时候,缓存也要跟着合并和丢弃,逻辑比贪心解码复杂不少。如果实现不当,很容易出现缓存和路径对不上的问题。
6. 从 KV Cache 延伸出去的几个思考
KV Cache 看起来只是“存一下 K 和 V”,但它其实是整个大模型推理优化的一个缩影。它把“计算换存储”这个经典权衡摆在了台面上:不做缓存,算力扛不住;做了缓存,显存扛不住。所有的优化手段,本质上都是在找这两者之间的平衡点。
我自己的体会是,理解 KV Cache 最好的方式不是背公式,而是亲手写一遍带缓存的注意力,然后跑一个长序列生成,看着显存曲线和延迟曲线,你自然就明白每一步在发生什么。纸上推一百遍 O(n²) 和 O(n),不如实际 profile 一次。
另外,KV Cache 的设计也影响模型架构的选择。为什么现在新模型越来越多用 GQA、MQA,为什么 RoPE 成了标配,为什么长上下文模型要专门设计注意力模式——这些决策背后都有 KV Cache 的影子。从这个角度说,搞懂 KV Cache,不只是搞懂一个优化技巧,而是搞懂了现代大模型推理的一条主线。
如果你正在调推理服务,建议先把每 token 的缓存字节数算清楚,再对照你的显存和并发目标,看看缺口有多大。这个数字一出来,该量化还是该换注意力方案,方向基本就定了。