KV Cache 注意力公式

计算公式

注意力权重统一记为 α,Attention 输出统一记为 o。o 经过多头合并、残差、FFN 和后续层后形成最终隐藏状态 h[N-1],用于预测第 N 个 Token。

   有 KV Cache:
   每一层只计算当前 Token 的 o[N-1],
   历史 Token 的 K/V 直接读取缓存。

   无 KV Cache:
   每一层重新处理完整前缀,
   从 o[1] 一直计算到 o[N-1],
   借此重建各层历史 K/V。
$$ \boldsymbol{\alpha}_{N-1} = \left[ \alpha_1,\alpha_2,\ldots,\alpha_{N-1} \right] = \operatorname{softmax} \left( \frac{ \left[ q_{N-1}k_1^T\;\; q_{N-1}k_2^T\;\; \cdots\;\; q_{N-1}k_{N-1}^T \right] }{\sqrt{d_h}} \right) $$
$$ o_{N-1} = \boldsymbol{\alpha}_{N-1} \begin{bmatrix} v_1\\ v_2\\ \vdots\\ v_{N-1} \end{bmatrix} = \alpha_1v_1 + \alpha_2v_2 + \cdots + \alpha_{N-1}v_{N-1} $$
$$ o_{N-1} = \frac{ \underbrace{ \exp\!\left(\frac{q_{N-1}k_1^T}{\sqrt{d_h}}\right)v_1 + \exp\!\left(\frac{q_{N-1}k_2^T}{\sqrt{d_h}}\right)v_2 + \cdots + \exp\!\left(\frac{q_{N-1}k_{N-2}^T}{\sqrt{d_h}}\right)v_{N-2} }_{k_1,\ldots,k_{N-2},v_1,\ldots,v_{N-2}\text{ 已缓存,查询分数本轮计算}} + \underbrace{ \exp\!\left(\frac{q_{N-1}k_{N-1}^T}{\sqrt{d_h}}\right)v_{N-1} }_{q_{N-1},k_{N-1},v_{N-1}\text{ 本轮新算}} }{ \underbrace{ \exp\!\left(\frac{q_{N-1}k_1^T}{\sqrt{d_h}}\right) + \exp\!\left(\frac{q_{N-1}k_2^T}{\sqrt{d_h}}\right) + \cdots + \exp\!\left(\frac{q_{N-1}k_{N-2}^T}{\sqrt{d_h}}\right) }_{k_1,\ldots,k_{N-2}\text{ 来自缓存,查询分数本轮计算}} + \underbrace{ \exp\!\left(\frac{q_{N-1}k_{N-1}^T}{\sqrt{d_h}}\right) }_{q_{N-1},k_{N-1}\text{ 本轮新算}} } $$

生成一个 Token 的单步复杂度

这里的“单步”指:根据前 N-1 个 Token,执行一次前向传播并预测第 N 个 Token。

无 KV Cache

模型重新处理完整前缀。每层、每个 Head 都要重新计算从 α[1] 到 α[N-1] 的全部注意力权重。应用 mask 后,有效 Q/K 点积数量为:

$$ 1+2+\cdots+(N-1) = \frac{N(N-1)}{2} = O(N^2) $$

用这些权重加权 V 的计算量同样为 O(N²),因此生成一个 Token 的单步注意力计算量为 O(N²)。

有 KV Cache

历史 Token 在每层产生的 K/V 已缓存。本轮只计算最后一个 Query 对全部 N-1 个 Key 的权重向量 α[N-1],再用它加权 N-1 个 Value:

$$ (N-1)+(N-1) = 2(N-1) = O(N) $$
根据前 N-1 个 Token 预测第 N 个 Token 单步注意力计算量
无 KV Cache,重算完整注意力矩阵 O(N²)
有 KV Cache,只算最后一行 O(N)

如果是生成1-L个token,那就是再翻一个数量级

KV Caching in LLMs, explained visually

显存分析

$$ \text{Memory}_{KV} \approx 2 \times b_{kv} \times L \times B \times S \times H \times \frac{N_{kv}}{N_{attn}} $$

其中:

  • $2$:同时缓存 Key 和 Value 矩阵。
  • $b_{kv}$:数据精度(Bytes),如 FP16 为 2。
  • $L$:模型层数 (Layers)。
  • $B$:并发请求数 (Batch Size)。
  • $S$:每个请求的平均序列长度(Prompt + 已生成 Token)。
  • $H$:隐藏层维度 (Hidden Size)。
  • $N_{kv}$:KV Head 的数量(GQA/MQA 中的分组数)。
  • $N_{attn}$:Query Head 的数量(总注意力头数)。当使用 MHA 时 $N_{kv} = N_{attn}$,系数为 1。

kvcache 的显存占用和 BatchSize 与 Token 序列长度成正比

hello world ! 的预测过程

假设 hello、world、! 分别是一个 Token。

用 hello 预测 world

模型当前只有 hello。单个 Head 的注意力权重和输出为:

$$ \alpha_{hello} = \operatorname{softmax} \left( \frac{q_{hello}k_{hello}^{T}}{\sqrt{d_h}} \right) =1, \qquad o_{hello}=v_{hello} $$

经过全部 Transformer 层后得到最后一个位置的隐藏状态,再由 LM Head 输出词表概率:

$$ P(world\mid hello) = \operatorname{softmax} \left( \operatorname{Norm}\!\left(h_{hello}^{(L)}\right)W_{vocab} \right) $$

本轮采样得到 world。

用 hello world 预测 !

hello 的 K/V 已缓存。本轮输入 world,每层只为它新计算 q_world、k_world 和 v_world:

$$ \left[ \alpha_{hello},\alpha_{world} \right] = \operatorname{softmax} \left( \frac{ \left[ q_{world}k_{hello}^{T}, q_{world}k_{world}^{T} \right] }{\sqrt{d_h}} \right) $$
$$ o_{world} = \alpha_{hello}v_{hello} + \alpha_{world}v_{world} $$

经过全部 Transformer 层后,world 位置的最终隐藏状态包含对 hello 和 world 的多轮融合。LM Head 据此计算:

$$ P(!\mid hello,world) = \operatorname{softmax} \left( \operatorname{Norm}\!\left(h_{world}^{(L)}\right)W_{vocab} \right) $$

本轮采样得到 !,序列成为 hello world !。! 的 Q/K/V 在下一轮计算,用于预测后面的 Token。