计算公式
注意力权重统一记为 α,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。