Menu

KV Cache:原理与显存占用计算

post on 12 Aug 2026 about 2638words require 9min
CC BY 4.0 (除特别声明或转载文章外)
如果这些文字帮助到你,可以请我喝一杯咖啡~

KV Cache 是大模型自回归推理中的一种优化:将各层已经计算过的 Key 和 Value 保存下来,在后续生成中复用,以额外的存储空间换取更少的重复计算。

本文以使用因果自注意力的 Decoder-only Transformer 为例,先回顾推理流程,再解释缓存原理,最后计算它的显存占用。

回顾 Transformer 的推理流程

输入向量经过权重矩阵投影,得到 Query(Q)、Key(K)和 Value(V)。Q 与 K 的转置相乘,得到注意力分数;经过缩放、掩码和 softmax 后,再与 V 相乘,得到融合上下文信息的输出。

对于一个注意力头,可以写成:

\[\operatorname{Attention}(Q,K,V) =\operatorname{softmax}\left(\frac{QK^\top}{\sqrt{d_k}}+M\right)V\]

其中,$d_k$ 是 Key 的维度,$M$ 是因果掩码,用来阻止当前位置看到未来的 token。因此,Q @ K.T @ V 只能作为粗略记法,实际计算还包含缩放与 softmax 等步骤。

KV Cache 原理示意图

Transformer 推理过程示意图之一

Transformer 推理过程示意图之二

Transformer 推理过程示意图之三

经过多层 Transformer 后,模型通过 lm_head 将隐藏状态投影到词表,得到 logits,再据此选择下一个 token。生成时,我们只需要最后一个位置的输出。例如,为了直观说明,假设“我”“爱”“吃”各对应一个 token:输入“我爱吃”后,最后一个位置的输出可用于预测接下来是否是“苹”。

这里省略了位置编码、残差连接、归一化和前馈网络等细节,重点关注注意力计算。

为什么历史 K、V 可以缓存?

自回归生成每次预测一个 token,再将它接到已有序列后面,继续预测下一个 token。最直接的实现会反复将整个序列输入模型,导致历史 token 的计算被重复执行。

自回归生成中的重复计算

在因果自注意力中,历史位置只能关注自身及更早的位置。追加新 token 不会改变历史位置可见的上下文,因此,在模型参数、前缀及位置处理保持一致的推理过程中,历史 token 在各层对应的 K、V 可以复用。

可复用的历史 Key 与 Value

需要注意,缓存的是每一层各自的 K、V 张量,而不是模型的权重矩阵,也不是一份供所有层共用的 K、V。

KV Cache 常用于解码器的自回归推理;Encoder–Decoder 模型的解码器交叉注意力也可以缓存来自编码器输出的 K、V。本文只讨论前一种情况。

每一步如何使用 KV Cache?

推理可以分为两个阶段:

  1. Prefill(预填充):处理完整的输入提示词,建立各层的 K、V 缓存,并利用最后一个位置的输出预测第一个新 token。
  2. Decode(逐 token 解码):将刚生成的 token 输入模型,在每一层只计算这个新位置的 Q、K、V,将新 K、V 追加到该层缓存中,再计算新位置的注意力输出。

假设当前输入的新 token 位于第 $t$ 个位置,在某一层中:

\[K_{1:t}=[K_{1:t-1};k_t],\qquad V_{1:t}=[V_{1:t-1};v_t]\] \[o_t=\operatorname{softmax}\left(\frac{q_tK_{1:t}^{\top}}{\sqrt{d_k}}\right)V_{1:t}\]

这里的分号表示沿序列维度拼接。对单序列、无 padding 的逐 token 解码,新位置可以看到缓存中的全部历史位置和自身,因此上式没有显式写出掩码。实际批处理仍需正确处理 attention mask 和位置索引。

新 token 的 Q、K、V 与历史缓存拼接

为什么通常不缓存 Q? 因为预测接下来的 token 只需要当前新位置的 Query;历史位置的 Query 不再参与这个新位置的注意力计算。历史的 K、V 则仍需被当前 Query 访问。

KV Cache 的收益来自减少历史位置的重复计算。新 token 仍然需要与可见历史的 K、V 做注意力运算,所以在完整注意力下,上下文越长,每一步需要读取和计算的缓存也越多。

多轮对话与上下文长度

多轮对话通常会将之前的消息放进后续请求的上下文中。如果推理系统保留并支持复用完全匹配的前缀缓存,就可以在这个前缀基础上继续处理新增消息;但聊天记录的保留并不意味着底层 KV Cache 一定跨请求保留。

随着缓存中的 token 增多,存储占用通常也会增加。超过上下文上限后的处理取决于具体模型和应用:可能截断历史、总结历史,也可能拒绝超长输入,不能一概理解为自动保留最后一段窗口。

例如,如果最早的“请将以下内容翻译成英文”指令被移出了输入上下文,后续模型就无法再直接看到这条要求。这个现象与上下文管理有关,不是 KV Cache 本身会主动遗忘指令。

计算 KV Cache 的显存占用

对于各层结构相同、使用完整注意力的标准 KV Cache,其张量存储量可以估算为:

\[\text{KV Cache bytes}=2\times B\times T\times L\times H_{\mathrm{KV}}\times D\times S\]
符号 含义
$2$ K 和 V 两份张量
$B$ 批次大小,即同时缓存的序列数
$T$ 缓存的序列长度,包含已处理的提示词及生成 token
$L$ Transformer 层数
$H_{\mathrm{KV}}$ 每层的 KV 头数
$D$ 每个头的维度,假设 K、V 头维度相同
$S$ 每个缓存元素的字节数,例如 FP16 / BF16 为 2 字节

笔记中的 KV Cache 内存计算示意图

原笔记中的“多头数”需要具体区分:普通多头注意力(MHA)的 KV 头数等于 Query 头数;分组查询注意力(GQA)的 KV 头数更少;多查询注意力(MQA)只有一个 KV 头。因此,应代入模型实际的 KV 头数,而不是一律使用 Query 头数。

举例来说,假设批次大小为 1、缓存长度为 4096、层数为 32、KV 头数为 32、每头维度为 128,并以 FP16 保存缓存:

\[2\times1\times4096\times32\times32\times128\times2 =2{,}147{,}483{,}648\ \text{bytes}=2\ \text{GiB}\]

如果其他条件不变,KV 头数改为 8,则缓存量降至 512 MiB。

这只是 KV 张量本身的估算,不包含模型权重、临时激活、内存分配开销等。静态预分配、滑动窗口注意力、缓存量化等实现也会影响实际占用;模型权重的精度与缓存精度需要分别确认。

参考资料

Related posts

Loading comments...