概述
KV Cache(Key-Value Cache) 是自回归生成时的一种加速技术。
在 Transformer 解码器中,生成文本是一个 token 一个 token 进行的:
text
输入 prompt → 生成第 1 个 token
把第 1 个 token 拼到输入后面 → 生成第 2 个 token
把第 2 个 token 拼到输入后面 → 生成第 3 个 token
每一步,模型都要计算当前所有 token 的注意力。
如果每一步都重新计算整个序列的 Key 和 Value,会有大量重复计算。
KV Cache 的做法是:
把之前每一步算好的 Key 和 Value 缓存起来,生成新 token 时只计算新 token 的 Key 和 Value, 然后追加到缓存中。
注意力计算时,用新 token 的 Query 去和整个缓存的 Key 交互,再对缓存的 Value 加权求和。
这样避免了重复计算历史 token 的 K 和 V。
为什么需要 KV Cache?
没有 KV Cache 时
假设已经生成了 ttt 个 token,现在要生成第 t+1t+1t+1 个。
- 需要重新计算这 ttt 个 token 的 Q、K、V;
- 但其中前 ttt 个 token 的 K、V 在之前步骤已经算过;
- 重复计算浪费了大量算力。
有 KV Cache 时
- 缓存中已经有前 ttt 个 token 的 K 和 V;
- 只需要计算新 token 的 Q、K、V;
- 把新的 K、V 追加到缓存;
- 用新的 Q 和整个缓存的 K 做点积,对缓存的 V 加权求和。
作用:
- 大幅加速自回归生成;
- 训练时不使用,因为训练可以并行处理整个序列;
- 推理/生成阶段几乎必用。
例子
为了能手算,我们取一个极小的 Transformer 解码器:
text
d_model = 2
单头注意力
d_k = d_v = 2
W_Q = W_K = W_V = 单位阵
没有位置编码
因果掩码:第 t 步只能看 1..t
假设依次输入的 token 向量是:
text
x1 = [1, 0]
x2 = [0, 1]
x3 = [1, 1]
第 1 步:生成第 1 个 token
计算:
text
Q1 = x1 = [1, 0]
K1 = x1 = [1, 0]
V1 = x1 = [1, 0]
缓存:
text
K_cache = [[1, 0]]
V_cache = [[1, 0]]
注意力分数:只有一个 key,softmax 权重为 1。
输出:
text
Z1 = V1 = [1, 0]
第 2 步:生成第 2 个 token
新 token:
text
x2 = [0, 1]
只计算新 token 的 Q、K、V:
text
Q2 = [0, 1]
K2 = [0, 1]
V2 = [0, 1]
把 K2、V2 追加到缓存:
text
K_cache = [[1, 0],
[0, 1]]
V_cache = [[1, 0],
[0, 1]]
用 Q2 和整个 K_cache 计算注意力分数:
Q2⋅K1=0,1⋅1,0=0 Q2 \cdot K1 = 0, 1 \cdot 1, 0 = 0 Q2⋅K1=0,1⋅1,0=0
Q2⋅K2=0,1⋅0,1=1 Q2 \cdot K2 = 0, 1 \cdot 0, 1 = 1 Q2⋅K2=0,1⋅0,1=1
除以 dk=2≈1.414\sqrt{d_k} = \sqrt{2} \approx 1.414dk =2 ≈1.414:
为什么要除以dk\sqrt{d_k}dk
点积的结果会随着维度增大而变大,如果不缩放,softmax 的输入会很大,导致梯度接近 0,训练不稳定。
除以 dk\sqrt{d_k}dk 可以把点积的方差拉回到 1 附近,保持梯度稳定。
0/1.414=0 0/1.414 = 0 0/1.414=0
1/1.414≈0.7071 1/1.414 \approx 0.7071 1/1.414≈0.7071
text
scores = [0, 0.7071]
Softmax:
e0=1,e0.7071≈2.028e^0 = 1, \quad e^{0.7071} \approx 2.028e0=1,e0.7071≈2.028
分母:
1+2.028=3.0281 + 2.028 = 3.0281+2.028=3.028
权重:
text
w1 = 1 / 3.028 ≈ 0.3302
w2 = 2.028 / 3.028 ≈ 0.6698
输出:
Z2=0.3302⋅V1+0.6698⋅V2 Z2 = 0.3302 \cdot V1 + 0.6698 \cdot V2 Z2=0.3302⋅V1+0.6698⋅V2
=0.3302⋅1,0+0.6698⋅0,1 = 0.3302 \cdot 1, 0 + 0.6698 \cdot 0, 1 =0.3302⋅1,0+0.6698⋅0,1
=0.3302,0.6698 = 0.3302, 0.6698 =0.3302,0.6698
第 3 步:生成第 3 个 token
新 token:
text
x3 = [1, 1]
只计算新 token 的 Q、K、V:
text
Q3 = [1, 1]
K3 = [1, 1]
V3 = [1, 1]
追加到缓存:
text
K_cache = [[1, 0],
[0, 1],
[1, 1]]
V_cache = [[1, 0],
[0, 1],
[1, 1]]
用 Q3 和整个 K_cache 计算分数:
Q3⋅K1=1,1⋅1,0=1 Q3 \cdot K1 = 1, 1 \cdot 1, 0 = 1 Q3⋅K1=1,1⋅1,0=1
Q3⋅K2=1,1⋅0,1=1 Q3 \cdot K2 = 1, 1 \cdot 0, 1 = 1 Q3⋅K2=1,1⋅0,1=1
Q3⋅K3=1,1⋅1,1=2 Q3 \cdot K3 = 1, 1 \cdot 1, 1 = 2 Q3⋅K3=1,1⋅1,1=2
除以 2\sqrt{2}2 :
1/1.414≈0.7071 1/1.414 \approx 0.7071 1/1.414≈0.7071
1/1.414≈0.7071 1/1.414 \approx 0.7071 1/1.414≈0.7071
2/1.414=2/2=2≈1.4142 2/1.414 = 2/\sqrt{2} = \sqrt{2} \approx 1.4142 2/1.414=2/2 =2 ≈1.4142
text
scores = [0.7071, 0.7071, 1.4142]
Softmax:
e0.7071≈2.028,e1.4142≈4.113 e^{0.7071} \approx 2.028, \quad e^{1.4142} \approx 4.113 e0.7071≈2.028,e1.4142≈4.113
分母:
2.028+2.028+4.113=8.169 2.028 + 2.028 + 4.113 = 8.169 2.028+2.028+4.113=8.169
权重:
text
w1 = 2.028 / 8.169 ≈ 0.2483
w2 = 2.028 / 8.169 ≈ 0.2483
w3 = 4.113 / 8.169 ≈ 0.5034
输出:
Z3=0.2483⋅V1+0.2483⋅V2+0.5034⋅V3 Z3 = 0.2483 \cdot V1 + 0.2483 \cdot V2 + 0.5034 \cdot V3 Z3=0.2483⋅V1+0.2483⋅V2+0.5034⋅V3
=0.2483⋅1,0+0.2483⋅0,1+0.5034⋅1,1 = 0.2483 \cdot 1, 0 + 0.2483 \cdot 0, 1 + 0.5034 \cdot 1, 1 =0.2483⋅1,0+0.2483⋅0,1+0.5034⋅1,1
第一维:
0.2483×1+0.2483×0+0.5034×1=0.7517 0.2483 \times 1 + 0.2483 \times 0 + 0.5034 \times 1 = 0.7517 0.2483×1+0.2483×0+0.5034×1=0.7517
第二维:
0.2483×0+0.2483×1+0.5034×1=0.7517 0.2483 \times 0 + 0.2483 \times 1 + 0.5034 \times 1 = 0.7517 0.2483×0+0.2483×1+0.5034×1=0.7517
所以:
Z3=0.7517,0.7517 Z3 = 0.7517, 0.7517 Z3=0.7517,0.7517
对比:无 KV Cache vs 有 KV Cache
| 步骤 | 无 KV Cache | 有 KV Cache |
|---|---|---|
| 第 1 步 | 计算 K1,V1 | 计算 K1,V1,缓存 |
| 第 2 步 | 重新计算 K1,V1,再算 K2,V2 | 只算 K2,V2,追加缓存 |
| 第 3 步 | 重新计算 K1,V1,K2,V2,再算 K3,V3 | 只算 K3,V3,追加缓存 |
| 总计算量 | O(L2d)O(L^2d)O(L2d) | O(Ld)O(Ld)O(Ld) 用于新 K,V,注意力仍 O(L2)O(L^2)O(L2) |
| 内存 | 无额外缓存 | 需要缓存所有历史 K,V |
可以看到,KV Cache 把"重复计算历史 K、V"的部分省掉了,生成越长,节省越多。
KV Cache 的内存开销
缓存大小公式:
KV Cache 大小=2×batch_size×num_heads×seq_len×head_dim×dtype_bytes \text{KV Cache 大小} = 2 \times \text{batch\_size} \times \text{num\_heads} \times \text{seq\_len} \times \text{head\_dim} \times \text{dtype\_bytes} KV Cache 大小=2×batch_size×num_heads×seq_len×head_dim×dtype_bytes
其中:
- 2 表示 K 和 V;
- batch_size 是批大小;
- num_heads 是注意力头数;
- seq_len 是当前序列长度;
- head_dim 是每个头的维度;
- dtype_bytes 是数据类型字节数(如 fp16 为 2)。
例如:
text
batch_size = 1
num_heads = 1
seq_len = 3
head_dim = 2
fp16 = 2 bytes
缓存大小:
2×1×1×3×2×2=24 bytes 2 \times 1 \times 1 \times 3 \times 2 \times 2 = 24 \text{ bytes} 2×1×1×3×2×2=24 bytes
真实模型中,seq_len 可能上万,num_heads 几十,head_dim 上百,缓存会非常大,因此出现了 MQA、GQA、KV Cache 量化等优化。
总结
KV Cache 是什么?
自回归生成时,缓存历史 token 的 Key 和 Value,生成新 token 时只计算新 token 的 K、V,避免重复计算。
在 Transformer 中的作用:
加速推理生成,把每步重复计算历史 K、V 的开销省掉;训练时不用,推理时几乎必用。
数值例子核心结论:
text
第 1 步:缓存 K1=[1,0], V1=[1,0],输出 Z1=[1,0]
第 2 步:只算 K2=[0,1], V2=[0,1],输出 Z2=[0.3302, 0.6698]
第 3 步:只算 K3=[1,1], V3=[1,1],输出 Z3=[0.7517, 0.7517]
每一步都只计算新 token 的 K、V,然后和整个缓存做注意力,这就是 KV Cache 的核心机制。