大语言模型研究15——注意力机制优化之KV缓存

作者: 引线小白-本文永久链接:https://www.limoncc.com/post/64e24a5816e7035f/
知识共享许可协议: 本博客采用署名-非商业-禁止演绎4.0国际许可证

一、基础概念

什么是 BMM (Batch Matrix Multiplication)。BMM 是批矩阵乘法,指的是对同一批次中的多个矩阵同时执行矩阵乘法。在 PyTorch 中,torch.matmul 或 @ 运算符在处理三维及以上张量时,遵循以下规则:把前 N−2 个维度当作“批次维度”,把最后 2 个维度当作“矩阵维度”进行乘法。它并不是一个独有的数学定义,而是深度学习框架为了高效处理“多组独立矩阵乘法”而提供的一种运算,通常对应 API 名称 bmm 或 matmul 的批处理能力。它并不是一个独有的数学定义,而是深度学习框架为了高效处理“多组独立矩阵乘法”而提供的一种运算,通常对应 API 名称 bmmmatmul 的批处理能力。下面重点讲它的形状约定和在不同场景下的行为。

1.1、严格 BMM 约定

PyTorch 中的 torch.bmm(input, mat2) 是最经典的 BMM 约定,规则非常严格:

  • 输入必须是两个 3 维张量
  • 两个张量的第一维(batch 维)必须相等。
  • 不支持广播,如果 batch 大小不同会直接报错。
1
2
3
4
5
# PyTorch 示例
import torch
a = torch.randn(10, 3, 4) # (B=10, M=3, K=4)
b = torch.randn(10, 4, 5) # (B=10, K=4, N=5)
c = torch.bmm(a, b) # 输出形状 (10, 3, 5)
1.2、广义批量矩阵乘法(如 torch.matmul / tf.matmul / numpy.matmul

现代框架中的 matmul@ 运算符支持任意多维张量,把最后两个维度作为矩阵,前面的所有维度都视为批维度,并遵循广播(broadcasting)规则。约定:

  • 如果张量维度 >2,则前导维度(除去最后两维以外的部分)会进行广播。
  • 矩阵维度遵循:(…, M, K) × (…, K, N) → (…, M, N)
1
2
3
4
# 示例:PyTorch 中的 matmul
a = torch.randn(2, 1, 3, 4) # 前导形状 (2,1) 矩阵 (3,4)
b = torch.randn(1, 5, 4, 2) # 前导形状 (1,5) 矩阵 (4,2)
c = a @ b # 广播后前导 (2,5),输出形状 (2,5,3,2)

这里 batch 维不必相等,只要可广播即可。这正是 Transformer 多头注意力中同时处理多个 head 和 batch 的基础。

1.3、应用场景
  • 全连接层的批量计算
    输入 (B, features) 可以看作 (B, 1, K),权重 (K, N) 拓展为 (1, K, N),用 BMM 得到 (B, 1, N),等价于 linear 的批处理。

  • 注意力机制中的批量点积
    Query: (B, num_heads, seq_len, d_k)
    Key 转置: (B, num_heads, d_k, seq_len)
    两者使用 matmul 直接得到注意力分数 (B, num_heads, seq_len, seq_len),同时利用了 head 和 batch 的双重批量。

  • 图神经网络中的边特征聚合
    邻接矩阵的批次运算,多个图样本同时进行消息传递。

1.4、与爱因斯坦求和(einsum)的关系

BMM 用 einsum 可以表示为:

  • torch.bmm(a, b) 对应 torch.einsum(‘bmk,bkn->bmn’, a, b)
  • 广义批量乘法则是 ‘…mk,…kn->…mn’,省略号代表广播的前导维度。

理解 einsum 后,BMM 的约定其实就浓缩为:共享 batch 下标,矩阵内积求和

1.5、总结

BMM 的核心约定就是:把独立的一批矩阵乘法打包,用统一的操作并行计算,要求(或通过广播使)除最后两个矩阵维度以外的所有前导维度对齐,内部仍遵守二维矩阵乘法规则。如果你在代码中遇到 bmmmatmul,根据是否需要广播和输入维度数来选择即可。

二、KV Cache 原理

在大型语言模型(LLM)的自回归生成过程中,推理速度面临极大的挑战。核心瓶颈在于:每生成一个新 Token,都需要对所有历史 Token 进行 Attention 计算。如果不加优化,生成长度为 $N$ 的序列,总计算量将是 $O(N^2)$。KV Cache 和 GQA (Grouped-Query Attention)是解决这一瓶颈的两大利器:前者通过空间换时间避免重复计算,后者通过结构优化大幅缩减缓存体积。

根据 Attention 机制,当前 Token 的 Query 会和所有历史 Token 的 Key 计算相似度,再和历史 Token 的 Value 加权求和。

  • 无 Cache 的痛点:生成第 $t$ 个 Token 时,我们需要把 $1$ 到 $t-1$ 的 Token 重新送入模型算出 $K_{1..t-1}$ 和 $V_{1..t-1}$。这导致历史 Token 被反复计算了 $N$ 次。
  • KV Cache 的核心思想:既然 Attention 计算只需要 $K$ 和 $V$,我们可以在生成第 $t-1$ 个 Token 时,把对应的 $K_{t-1}$ 和 $V_{t-1}$ 缓存下来。生成第 $t$ 个 Token 时,只需计算当前 Token 的 $Q_t, K_t, V_t$,然后将 $K_t, V_t$ 拼接到历史的 Cache 中即可。

推理的两个阶段

  1. Prefill 阶段(预填充):输入 Prompt,并行计算所有 Token 的 KV 并缓存。此阶段属于计算密集型
  2. Decode 阶段(解码):逐个生成 Token,每步读取历史 KV Cache,并追加当前步的 KV。此阶段属于显存访问密集型

KV Cache 并非没有代价,它消耗巨大的显存。对于模型层数 $L$,隐藏层维度 $d_{model}$,序列长度 $N$,精度为 $b$ 字节(如 FP16 为 2 字节):

$$\begin{align}
\text{KV Cache Size} = 2 \times L \times N \times d_{model} \times b
\end{align}$$

以 LLaMA-2 70B 为例:$L=80, d_{model}=8192, N=4096, b=2$,单条序列的 KV Cache 约需 10GB 显存!为了减少 KV Cache 体积,GQA 应运而生。

三、从 MHA 到 GQA

3.1、标准 Multi-Head Attention (MHA)

在 MHA 中,有 $h$ 个 Query 头,$h$ 个 Key 头,$h$ 个 Value 头。对于输入 $\bm{X} \in \mathbb{R}^{N \times d}$:

$$\begin{align}
\bm{Q} = \bm{X}\bm{W}_Q, \quad \bm{K} = \bm{X}\bm{W}_K, \quad \bm{V} = \bm{X}\bm{W}_V
\end{align}$$

其中 $\bm{W}_Q \in \mathbb{R}^{d \times (h \cdot d_k)}$, $\bm{W}_K \in \mathbb{R}^{d \times (h \cdot d_k)}$, $\bm{W}_V \in \mathbb{R}^{d \times (h \cdot d_v)}$。

Attention 计算公式为:

$$ \text{Attention}(\bm{Q}_i, \bm{K}_i, \bm{V}_i) = \mathrm{softmax}\left(\frac{\bm{Q}_i \bm{K}_i^T}{\sqrt{d_k}}\right) \bm{V}_i $$
KV Cache 体积:与头数 $h$ 成正比。

3.2、Multi-Query Attention (MQA)

MQA 极端地让所有 Query 头共享1个 Key 和 Value 头。

$$\begin{align}
h_K = h_V = 1
\end{align}$$

这极大地节省了 Cache,但由于参数量骤减,可能导致模型质量下降。

3.3、Grouped-Query Attention (GQA) —— 完美的折中

GQA 将 $h$ 个 Query 头分为 $g$ 个组,每个组共享1个 Key 头和1个 Value 头。定义每组包含 $m = h/g$ 个连续的 Query 头。则,对于索引为 $i\in [0, h-1]$ 的 Query 头,它属于第 $ \lfloor i/m \rfloor $ 组,因此使用的 KV 头索引也是这个组号:

$$\begin{align}
j = \left\lfloor \frac{i}{m} \right\rfloor = \left\lfloor \frac{i \cdot g}{h} \right\rfloor
\end{align}$$

即 Query 头 0 到 $m-1$ 共享 KV 头 0,Query 头 $m$ 到 $2m-1$ 共享 KV 头 1,依此类推。

$$\begin{align}
\text{Attention}(\bm{Q}_i, \bm{K}_j, \bm{V}_j) = \mathrm{softmax}\left(\frac{\bm{Q}_i \bm{K}_j^T}{\sqrt{d_k}}\right) \bm{V}_j
\end{align}$$

KV Cache 体积:缩小为 MHA 的 $\frac{g}{h}$ 倍!
例如 LLaMA-2 70B:$h=64, g=8$,KV Cache 体积缩小为原来的 1/8!

四、代码实现:带 KV Cache 的 GQA

4.1、层内处理

下面使用 PyTorch 从零实现一个带有 KV Cache 的 GQA 模块。核心逻辑拆解

  • 1、投影输出:计算当前输入的 $Q, K, V$。注意 $K, V$ 的形状是 [batch, seq_len, num_kv_heads, head_dim]
  • 2、Cache 拼接:将算出的 $K, V$ 与历史 Cache 拼接。
  • 3、GQA 扩展:将 $K, V$ 的 num_kv_heads 维度扩展为 num_q_heads,以便与 $Q$ 进行标准的 Batch Matrix Multiplication (BMM)。
  • 4、计算 Attention:标准的缩放点积注意力。
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
import torch
import torch.nn as nn
import math
class GQAttentionWithCache(nn.Module):
def __init__(self, d_model, num_q_heads, num_kv_heads):
super().__init__()
assert num_q_heads % num_kv_heads == 0, "num_q_heads must be divisible by num_kv_heads"

self.num_q_heads = num_q_heads
self.num_kv_heads = num_kv_heads
self.head_dim = d_model // num_q_heads
self.n_rep = num_q_heads // num_kv_heads # 每个 KV 头被复用的次数

# 线性投影层
self.W_Q = nn.Linear(d_model, num_q_heads * self.head_dim, bias=False)
self.W_K = nn.Linear(d_model, num_kv_heads * self.head_dim, bias=False)
self.W_V = nn.Linear(d_model, num_kv_heads * self.head_dim, bias=False)
self.W_O = nn.Linear(num_q_heads * self.head_dim, d_model, bias=False)

def repeat_kv(self, x):
"""将 KV 头扩展到与 Q 相同的头数以进行注意力计算
x shape: [batch, seq_len, num_kv_heads, head_dim]
return shape: [batch, seq_len, num_q_heads, head_dim]
"""
if self.n_rep == 1:
return x
batch, seq_len, num_kv_heads, head_dim = x.shape
# unsqueeze: 可以理解为升维——在张量的指定位置插入一个大小为 1 的新维度。
# squeeze: 删除张量中所有(或指定位置)大小为 1 的维度,从而降低维数。
# expand: 主要用于将形状中大小为 1 的维度扩展为更大的值,不分配新内存, 不占显存复制数据
x = x.unsqueeze(3).expand(batch, seq_len, num_kv_heads, self.n_rep, head_dim)
return x.reshape(batch, seq_len, num_kv_heads * self.n_rep, head_dim)
def forward(self, x, kv_cache=None, start_pos=0):
"""
x: 当前输入 Token 的 Embedding [batch_size, seq_len, d_model]
- Prefill 阶段: seq_len = prompt_length
- Decode 阶段: seq_len = 1
kv_cache: 元组
start_pos: 当前 Token 在序列中的起始位置
"""
batch_size, seq_len, _ = x.shape

# 1. 计算 Q, K, V
Q = self.W_Q(x) # [B, S, num_q_heads * head_dim]
K = self.W_K(x) # [B, S, num_kv_heads * head_dim]
V = self.W_V(x) # [B, S, num_kv_heads * head_dim]

# 重塑形状以便按头处理: [B, S, H, D] -> [B, H, S, D]
Q = Q.view(batch_size, seq_len, self.num_q_heads, self.head_dim).transpose(1, 2)
K = K.view(batch_size, seq_len, self.num_kv_heads, self.head_dim).transpose(1, 2)
V = V.view(batch_size, seq_len, self.num_kv_heads, self.head_dim).transpose(1, 2)

# 这里为简单记,省略位置编码
# cos, sin = position_embeddings
# Q, K = apply_rotary_pos_emb(Q, K, cos[:seq_len], sin[:seq_len])

# 2. 处理 KV Cache
if kv_cache is not None:
past_K, past_V = kv_cache
# past_K/V shape: [B, num_kv_heads, past_seq_len, head_dim]
# 将新的 K, V 拼接到历史 Cache 后面
K = torch.cat([past_K, K], dim=2)
V = torch.cat([past_V, V], dim=2)

# 更新 Cache 供下一步使用 (注意:实际框架中会使用就地操作或预分配张量优化)
new_kv_cache = (K, V)

# 3. GQA 的 K, V 扩展
# 此时 K, V shape: [B, num_kv_heads, full_seq_len, head_dim]
K_expanded = self.repeat_kv(K.transpose(1, 2)).transpose(1, 2) # -> [B, num_q_heads, full_seq_len, head_dim]
V_expanded = self.repeat_kv(V.transpose(1, 2)).transpose(1, 2) # -> [B, num_q_heads, full_seq_len, head_dim]

# 4. 计算缩放点积注意力
# Q shape: [B, num_q_heads, seq_len, head_dim]
# K_expanded shape: [B, num_q_heads, full_seq_len, head_dim]
scores = torch.matmul(Q, K_expanded.transpose(2, 3)) / math.sqrt(self.head_dim)

# 生成因果掩码,防止看到未来的 Token (仅 Prefill 需要,Decode 阶段 seq_len=1 自动满足)
if seq_len > 1:
mask = torch.tril(torch.ones(seq_len, K_expanded.size(2), device=x.device)).view(1, 1, seq_len, K_expanded.size(2))
scores = scores.masked_fill(mask == 0, float('-inf'))

attn_weights = torch.softmax(scores, dim=-1)

# context shape: [B, num_q_heads, seq_len, head_dim]
context = torch.matmul(attn_weights, V_expanded)

# 5. 合并多头并输出投影
context = context.transpose(1, 2).contiguous().view(batch_size, seq_len, -1)
output = self.W_O(context)

return output, new_kv_cache
# ================= 测试代码 =================
if __name__ == "__main__":
d_model = 512
num_q_heads = 8
num_kv_heads = 2 # GQA: 8个Q头共享2个KV头

layer = GQAttentionWithCache(d_model, num_q_heads, num_kv_heads)

# ================= 长序列对比 =================
print("\n" + "-" * 60)
print("长序列对比测试 (prompt_len=200, decode=50):")
print("-" * 60)
batch_size = 2
x_long = torch.randn(batch_size, 200, d_model)
x_long_decode = torch.randn(batch_size, 50, d_model)
full_long = torch.cat([x_long, x_long_decode], dim=1)

# 无缓存
torch.cpu.synchronize() if torch.cpu.is_available() else None
t0 = time.perf_counter()
for step in range(50):
cur_len = 200 + step + 1
out, _ = layer(full_long[:, :cur_len, :], kv_cache=None)
t1 = time.perf_counter()
long_no_cache = t1 - t0

# 有缓存
torch.cpu.synchronize() if torch.cpu.is_available() else None
t2 = time.perf_counter()
_, cache_long = layer(x_long, kv_cache=None)
for step in range(50):
x_step = full_long[:, 200 + step: 200 + step + 1, :]
out, cache_long = layer(x_step, kv_cache=cache_long)
t3 = time.perf_counter()
long_with_cache = t3 - t2

print(f" 无缓存: {long_no_cache:.4f}s | 有缓存: {long_with_cache:.4f}s | 加速比: {long_no_cache / long_with_cache:.2f}x")
4.2、层外处理

堆叠后就是这样了。

1
self.layers = nn.ModuleList([LLMBlock(l, config) for l in range(self.num_hidden_layers)])

在外层需要处理位置编码,我们使用past_key_values 存储所有层的KV缓存。

1
2
3
4
5
6
7
8
9
10
11
12
# 不存在,就先预留位置
past_key_values = past_key_values or [None] * len(self.layers)
# 计算已经 Prefill + Decode的token数
# past_key_values[0] 是第一层的 KV 缓存元组 (K, V),其中 K.shape = [batch_size, cached_seq_len, num_kv_heads, head_dim]
# K.shape[1] 就是已经缓存的 token 数量
start_pos = past_key_values[0][0].shape[1] if past_key_values[0] is not None else 0
# 计算位置编码后的向量需要cos和sin
# RoPE(x) = x * cos(θ) + rotate_half(x) * sin(θ)
# 预先计算出来
position_embeddings = (
self.freqs_cos[start_pos:start_pos + seq_length],
self.freqs_sin[start_pos:start_pos + seq_length])

然后向每层传递

1
2
3
4
5
6
7
8
9
10
kv_caches = []
for layer_idx, (layer, past_key_value) in enumerate(zip(self.layers, past_key_values)):
hidden_states, kv_cache = layer(
hidden_states,
position_embeddings,
past_key_value=past_key_value,
use_cache=use_cache,
attention_mask=attention_mask
)
kv_caches.append(kv_cache)

当然这里是极简实现,transformers库的实现要复杂的多,考虑到了更多工程问题。但是核心也就这些。在真实的生产环境中,除了 GQA,我们还需要解决 Cache 显存碎片问题。vLLM 提出的 PagedAttention 借鉴了操作系统的虚拟内存分页机制,将不连续的 KV Cache 显存块管理起来,结合 GQA,实现了极高的吞吐量。

五、前缀缓存(Prefix Caching)

KV Cache 如何通过空间换时间避免自回归生成时的重复计算。然而,在实际的大模型应用场景(如 RAG 检索增强生成、多轮对话、Agent 系统提示词)中,我们常常面临一个新的痛点:不同请求之间,往往包含大量完全相同的文本前缀。
如果每个请求都独立计算这部分相同前缀的 KV Cache,不仅浪费海量算力,更极大地增加了系统的首字响应时间(TTFT)。Prefix Caching(前缀缓存)正是为解决这一痛点而生。

5.1、Prefix Caching 核心思想

既然相同输入 Token 必然生成相同的 KV Cache,何不将其缓存起来供所有请求共享?
Prefix Caching 将 KV Cache 的生命周期从“单个请求的生存期”提升到了“全局跨请求的生存期”。当新请求到达时,系统首先检查其前缀是否已经被缓存:

  • 命中:直接从显存/内存中加载对应的 KV Cache,只需对新接入的 Token 计算 Prefill。
  • 未命中:正常计算,并将计算出的 KV Cache 写入缓存池。

物理实现的关键:PagedAttention(分页注意力)

在传统的连续 KV Cache 存储中,不同请求的 Cache 在显存中是分散且大小不一的,无法直接共享。Prefix Caching 的工业级实现(如 vLLM)必须依赖 PagedAttention

  • 将 KV Cache 切分为固定大小的 *Block(页),类似操作系统的虚拟内存分页。
  • 相同前缀的 KV Cache 指向同一组物理 Block。
  • 采用 Copy-on-Write(写时复制)机制:当请求 B 在前缀后生成新 Token 时,新 Token 的 KV Cache 会被写入新分配的 Block,而不会覆盖共享的前缀 Block。

代码实现:带 Prefix Caching 的 GQA
为了直观展示原理,以下代码不涉及复杂的底层显存分页管理,而是用 PyTorch 和字典模拟 Prefix Cache 的逻辑匹配与复用过程
我们复用上一讲的 GQAttentionWithCache,并在此基础上构建一个 PrefixCacheManager

5.2、PyTorch 模拟代码
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
import torch
import torch.nn as nn
import hashlib
# 复用上一讲的 GQA 模块 (此处省略详细定义,假设已导入)
# from previous_lecture import GQAttentionWithCache
class GQAttentionWithCache(nn.Module):
# ... (此处为节省篇幅,包含上一讲的完整实现) ...
pass
class PrefixCacheManager:
def __init__(self):
# 使用字典存储: {prefix_hash: (K_cache, V_cache)}
self.cache_pool = {}

def _hash_prefix(self, token_ids):
"""将 Token ID 序列哈希,作为前缀的唯一指纹"""
hasher = hashlib.sha256()
hasher.update(token_ids.cpu().numpy().tobytes())
return hasher.hexdigest()

def get_cache(self, token_ids):
"""查询前缀缓存"""
prefix_hash = self._hash_prefix(token_ids)
return self.cache_pool.get(prefix_hash, None)

def set_cache(self, token_ids, kv_cache):
"""写入前缀缓存"""
prefix_hash = self._hash_prefix(token_ids)
# 注意:实际工程中需要深度复制或引用计数,这里简化为直接存储
self.cache_pool[prefix_hash] = kv_cache
class LLMWithPrefixCaching(nn.Module):
def __init__(self, d_model, num_q_heads, num_kv_heads, num_layers=2):
super().__init__()
self.num_layers = num_layers
self.embedding = nn.Embedding(1000, d_model) # 假设词表1000
# 每一层都有一个独立的 Attention 和 独立的 Prefix Cache Manager
self.attn_layers = nn.ModuleList([
GQAttentionWithCache(d_model, num_q_heads, num_kv_heads)
for _ in range(num_layers)
])
self.cache_managers = [PrefixCacheManager() for _ in range(num_layers)]

def forward(self, input_ids, prefix_len=None):
"""
input_ids: [batch_size, seq_len] Token ID 序列
prefix_len: 标明前多少个 Token 属于可共享的前缀
"""
batch_size, seq_len = input_ids.shape

# 1. 切分前缀和新增部分
if prefix_len is not None and prefix_len > 0:
prefix_ids = input_ids[:, :prefix_len]
new_ids = input_ids[:, prefix_len:]
else:
prefix_ids = None
new_ids = input_ids

x = self.embedding(input_ids)
x_new = self.embedding(new_ids) if prefix_len is not None else x

layer_kv_caches = [] # 用于保存当前请求的完整 KV Cache (包含前缀+新增)
# 2. 逐层计算
for layer_idx, attn_layer in enumerate(self.attn_layers):
manager = self.cache_managers[layer_idx]

# --- 前缀缓存逻辑 ---
if prefix_len is not None and prefix_len > 0:
# 查询缓存
cached_kv = manager.get_cache(prefix_ids[0]) # 假设 Batch 内前缀一致

if cached_kv is not None:
# 缓存命中!
print(f"Layer {layer_idx}: Prefix Cache HIT! Skipping {prefix_len} tokens prefill.")
past_kv = cached_kv
else:
# 缓存未命中,计算前缀并存入缓存
print(f"Layer {layer_idx}: Prefix Cache MISS! Calculating prefix prefill.")
x_prefix = self.embedding(prefix_ids)
_, past_kv = attn_layer(x_prefix, kv_cache=None)
manager.set_cache(prefix_ids[0], past_kv)

# --- 新增部分计算 ---
# 将前缀的 KV 作为历史 Cache 传入,仅对 new_ids 计算 Prefill
x_new_out, updated_kv = attn_layer(x_new, kv_cache=past_kv)

# 更新当前请求的完整 KV (前缀 + 新增)
layer_kv_caches.append(updated_kv)

# 为下一层准备输入 (实际需要残差连接和 FFN,此处简化)
x_new = x_new_out
else:
# 无前缀,正常计算
x_out, kv = attn_layer(x, kv_cache=None)
layer_kv_caches.append(kv)
x_new = x_out

return x_new, layer_kv_caches
# ================= 测试代码 =================
if __name__ == "__main__":
d_model = 512
num_q_heads = 8
num_kv_heads = 2

model = LLMWithPrefixCaching(d_model, num_q_heads, num_kv_heads, num_layers=2)

# 场景:RAG 应用,System Prompt 有 10 个 Token
prefix_tokens = torch.randint(0, 1000, (1, 10))

# 用户 A 的提问:3 个 Token
query_A = torch.randint(0, 1000, (1, 3))
input_A = torch.cat([prefix_tokens, query_A], dim=1)

# 用户 B 的提问:4 个 Token
query_B = torch.randint(0, 1000, (1, 4))
input_B = torch.cat([prefix_tokens, query_B], dim=1)

print("=== 处理请求 A ===")
out_A, _ = model(input_A, prefix_len=10)

print("\n=== 处理请求 B (共享相同前缀) ===")
out_B, _ = model(input_B, prefix_len=10)

核心逻辑拆解:

  1. 哈希匹配:PrefixCacheManager 对 Token ID 序列求 SHA-256 哈希。只要用户的 Token 完全相同,哈希值就一致。
  2. MISS 分支:请求 A 首次到达,缓存为空。模型被迫对前 10 个 Token 做 Prefill,然后将计算出的 past_kv 存入 Manager。
  3. HIT 分支:请求 B 到达,发现前缀哈希命中。直接取出 past_kv,将其作为 kv_cache 参数传入 Attention,仅对后面的 3~4 个新 Token 执行 Prefill
5.3、工程实践:从逻辑到物理的跨越

上述代码演示了逻辑原理,但在真实的生产环境(如高并发服务)中,存在极大的工程挑战:

显存碎片与 PagedAttention

真实场景中,前缀长度千变万化。如果为每个前缀分配连续的 Tensor 显存,显存会迅速被碎片化,导致 OOM。
解法:vLLM 引入了 PagedAttention。将 KV Cache 分割为固定大小的 Block(如 16 个 Token 一个 Block)。前缀缓存以 Block 为单位存储。请求 B 命中前缀时,只需在页表中映射指向这些物理 Block,无需拷贝数据。

写时复制

请求 B 在前缀之后生成了新 Token,其 KV Cache 需要追加。由于前缀 Block 是共享的,绝对不能直接修改。
解法:新增的 KV 写入新分配的 Block 中,逻辑上通过链表/页表将其与共享的前缀 Block 串联起来,形成完整的 KV Cache 逻辑视图。

驱逐策略

显存有限,不可能缓存所有历史前缀。当显存不足时,需要淘汰旧缓存。
解法:类似操作系统的 LRU(最近最少使用)策略。vLLM 的 PrefixCaching 调度器会监控 Block 的引用计数和时间戳,优先淘汰没有请求在使用且最久未访问的前缀 Block。

RadixAttention (SGLang)

相比于单纯的前缀匹配,SGLang 提出了更激进的 RadixAttention(基数树注意力)。它将所有请求的 Token 序列在一棵 Radix Tree 上进行前缀匹配。不仅系统提示词可以共享,多轮对话的历史记录、Agent 中间步骤的公共子序列,都能在树形结构中找到最长公共前缀并复用 KV Cache。


版权声明
引线小白创作并维护的柠檬CC博客采用署名-非商业-禁止演绎4.0国际许可证。
本文首发于柠檬CC [ https://www.limoncc.com ] , 版权所有、侵权必究。
本文永久链接https://www.limoncc.com/post/64e24a5816e7035f/
如果您需要引用本文,请参考:
引线小白. (May. 19, 2026). 《大语言模型研究15——注意力机制优化之KV缓存》[Blog post]. Retrieved from https://www.limoncc.com/post/64e24a5816e7035f
@online{limoncc-64e24a5816e7035f,
title={大语言模型研究15——注意力机制优化之KV缓存},
author={引线小白},
year={2026},
month={May},
date={19},
url={\url{https://www.limoncc.com/post/64e24a5816e7035f}},
}

'