KV Cache

概述

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 的核心机制。

相关推荐
H.莓飛1 小时前
【数据结构】二叉树
linux·数据结构·算法
Carl_奕然1 小时前
【智能体】Loop 的四种设计模式之:Hill Climbing Loop(2026 最新版)
人工智能·python·设计模式
Rocky Ding*1 小时前
深度解析LlamaGen核心基础知识
论文阅读·人工智能·深度学习·机器学习·aigc·ai-native·llamagen
hai3152475431 小时前
语言学矩阵原理
人工智能·机器学习·矩阵
Omics Pro1 小时前
~30,000+引用!理论2005,R包2008,多组学集成AI增强
开发语言·数据库·人工智能·算法·机器学习·自然语言处理·r语言
ZzT1 小时前
拆开 Claude Code、Codex 等 11 个 coding agent:没有一个用 LangChain,也没有一个用向量检索代码
人工智能·ai编程·claude
荣合技术服务1 小时前
VMware ESXi 虚拟化平台服务器虚拟机数据恢复服务
linux·运维·服务器
Kstheme1 小时前
受 Karpathy ASD-STE100的启发,我把论文讲解做成了一个开源 Skill
人工智能
EatFan1 小时前
MCP 从概念到落地:Java(Spring AI Alibaba)与 .NET 双栈接入实操对比
java·人工智能·后端·spring·.net·java后端·mcp