GQA 与 KV 缓存(Grouped-Query Attention & KV Cache)

一句话本质 :GQA 通过让多个 Query 头共享少量 KV 头,将推理显存瓶颈压缩至原来的 1/8 而几乎不损质量;KV Cache 则利用因果掩码下历史 K/V 不变的特性缓存复用,把自回归生成从 O(n3)O(n^3) O(n3) 降到 O(n2)O(n^2) O(n2)------两者共同构成现代大模型高并发、低延迟推理的基石


1. GQA 为什么需要

1.1 MHA 的 KV Cache 显存瓶颈

标准多头注意力(MHA)中,每个 Query 头独占一组 Key/Value 头。推理时,历史 token 的 K/V 必须缓存(详见 §5),KV Cache 的显存占用为:
KV_bytes=2×L×n×nkv_heads×dhead×precision_bytes\text{KV\bytes} = 2 \times L \times n \times n{\text{kv\heads}} \times d{\text{head}} \times \text{precision\_bytes} KV_bytes=2×L×n×nkv_heads×dhead×precision_bytes

其中 nkv_heads n_{\text{kv\heads}} nkv_heads 是 KV 头数(MHA 中等于注意力头数 hh h;GQA/MQA 更少)。以 LLaMA-2 70B 若采用 MHA 为例( nkv_heads=h=64 n{\text{kv\heads}}=h=64 nkv_heads=h=64, dhead=128 d{\text{head}}=128 dhead=128, L=80L=80 L=80, FP16),单请求 4096 token 上下文
2×80×4096×64×128×2=10.74 GB2 \times 80 \times 4096 \times 64 \times 128 \times 2 = 10.74 \text{ GB} 2×80×4096×64×128×2=10.74 GB

并发 50 个请求就需要 ~537 GB 显存------远超模型权重本身(70B × 2B ≈ 140 GB)。KV Cache 才是推理阶段真正的显存大户

1.2 核心洞察:多样性源于 Q 而非 K/V

注意力公式 Output=softmax(QK⊤/ dk )⋅V\text{Output} = \text{softmax}(QK^\top / \sqrt{d_k}) \cdot V Output=softmax(QK⊤/dk )⋅V 中,Q 决定「看哪里」(注意力权重分布),V 决定「看到什么」(被提取的内容) 。实验发现:不同 Q 头学到的 K/V 表示高度冗余 ------同一组内的 Q 头虽然各自关注不同的语义模式(语法、指代、位置......),但它们需要检索的底层信息(由 K/V 编码)其实大同小异。

这意味着:不需要为每个 Q 头都复制一份 K/V,多个 Q 头完全可以共享同一组 K/V 头。这就是 GQA(Grouped-Query Attention)的出发点。


2. MHA / MQA / GQA 三方案对比

2.1 头分配关系

方案 Query 头数 KV 头数 分配方式 特点
MHA 多头注意力 hh h hh h 每个 Q 头独占 1 组 K/V 表达力最强、显存最大
MQA 多查询注意力 hh h 11 1 所有 Q 头共享 1 组 K/V 显存最省、质量有损
GQA 分组查询注意力 hh h gg g( 1<g<h1 < g < h 1<g<h) h/gh/g h/g 个 Q 头共享 1 组 K/V 折中,主流选择

头分配示意图( h=8h=8 h=8):

less 复制代码
标准 MHA (g=8):              GQA (g=2):                  MQA (g=1):
每个 Q 头独占 K/V            每 4 个 Q 头共享 1 组 K/V    所有 Q 头共享 1 组 K/V

Q0 ─ K0/V0                  Q0 ─┐                        Q0 ─┐
Q1 ─ K1/V1                  Q1 ─┤                        Q1 ─┤
Q2 ─ K2/V2                  Q2 ─┼─ K0/V0                 Q2 ─┤
Q3 ─ K3/V3                  Q3 ─┘                        Q3 ─┼─ K0/V0
Q4 ─ K4/V4                  Q4 ─┐                        Q4 ─┤
Q5 ─ K5/V5                  Q5 ─┤                        Q5 ─┤
Q6 ─ K6/V6                  Q6 ─┼─ K1/V1                 Q6 ─┤
Q7 ─ K7/V7                  Q7 ─┘                        Q7 ─┘

KV 头数 = 8                 KV 头数 = 2                  KV 头数 = 1

2.2 真实模型配置表

模型 Q 头数 hh h KV 头数 gg g 每组 Q 数 dmodel d_{\text{model}} dmodel 层数 LL L
LLaMA-2 70B 64 8 8 8192 80
LLaMA-3 8B 32 8 4 4096 32
LLaMA-3 70B 64 8 8 8192 80
Qwen2 72B 64 8 8 8192 80
Mistral 7B 32 8 4 4096 32
Gemma 2 27B 32 16 2 3584 46

趋势:KV 头数通常取 8(或 4/16),是「质量--显存」帕累托前沿的经验最优区间。GQA 已成为当前主流大模型的事实标准。


3. GQA-8 深入论证

以 LLaMA-2 70B(GQA-8, h=64h=64 h=64, g=8g=8 g=8, dk=128 d_k=128 dk=128, dmodel=8192 d_{\text{model}}=8192 dmodel=8192)为蓝本,从六个维度论证 GQA-8 的合理性。

3.1 维度一:多样性来自 Q,不来自 K/V

同一组内(共享 K0/V0 K_0 / V_0 K0/V0),8 个 Q 头的输出:
Oi=softmax ⁣ ( Qi⋅K0⊤ dk ) ⋅V0,i=0,1,...,7 O_i = \text{softmax}\!\left(\frac{Q_i \cdot K_0^\top}{\sqrt{d_k}}\right) \cdot V_0, \quad i = 0, 1, \ldots, 7 Oi=softmax(dk Qi⋅K0⊤)⋅V0,i=0,1,...,7

K/V 相同,但 8 个 Q 头产出 8 个完全不同的结果 ------因为不同的 Qi Q_i Qi 产生不同的注意力权重分布。

类比:8 个人去同一个图书馆 (同一份 K/V),但每人带着不同的问题 (不同的 Q):有人找语法线索、有人找指代关系、有人找主题词。图书馆不需要 64 套藏书------一套完整的藏书 + 64 种不同的检索策略 = 64 种不同的答案。多样性来自「提问角度」,不来自「复制图书馆」。

3.2 维度二:Q 和 K/V 的角色本质不同

Q(查询) K(键) V(值)
角色 需求侧:「我要找什么」 索引侧:「我能被怎么匹配」 供给侧:「选中我后给什么」
多样性需求 :不同头要问不同问题 :一个 token 的可匹配属性有限 :一个 token 的实际内容就那么多
类比 64 种不同读者 8 种编目方式 8 种内容摘要

K 是「标签/索引」:一个 token 能被匹配的方式是有限的------语法角色、语义类别、位置关系、指代属性、主题关联、情感色彩、局部搭配、长程依赖------8 个 KV 头已经能编码这些主要维度。再多就是冗余,你不会因为给一本书贴了 64 张标签就让它变得更好找。

V 是「内容载体」:一个 token 实际能贡献的信息量有限(它就是一个词的语义)。8 个 128 维向量(共 1024 维)已足以编码一个 token 的完整语义内容。给同一份内容做 64 份拷贝,不会让它「内容更丰富」。

3.3 维度三:表达力无损

MHA 的注意力输出:
OiMHA=softmax ⁣ ( Qi⋅Ki⊤ d ) ⋅Vi每个头有自己的 Ki,Vi O_i^{\text{MHA}} = \text{softmax}\!\left(\frac{Q_i \cdot K_i^\top}{\sqrt{d}}\right) \cdot V_i \quad \text{每个头有自己的 } K_i, V_i OiMHA=softmax(d Qi⋅Ki⊤)⋅Vi每个头有自己的 Ki,Vi

GQA-8 的注意力输出:
OiGQA=softmax ⁣ ( Qi⋅ K⌊i/8⌋⊤ d ) ⋅ V⌊i/8⌋ 同组共享 K,V O_i^{\text{GQA}} = \text{softmax}\!\left(\frac{Q_i \cdot K_{\lfloor i/8 \rfloor}^\top}{\sqrt{d}}\right) \cdot V_{\lfloor i/8 \rfloor} \quad \text{同组共享 } K, V OiGQA=softmax(d Qi⋅K⌊i/8⌋⊤)⋅V⌊i/8⌋同组共享 K,V

关键观察 :即使 K/V 相同,只要 Qi Q_i Qi 不同, softmax(Qi⋅K⊤) \text{softmax}(Q_i \cdot K^\top) softmax(Qi⋅K⊤) 就不同,加权求和 Ai⋅V A_i \cdot V Ai⋅V 的结果就不同。64 个 Q 头仍然产出 64 个不同的输出向量,最终拼接后同样是 8192 维------表达力并未因 K/V 共享而塌缩。

真正的差异在于:MHA 中每个头可以从「专属 V 空间」取信息,而 GQA 中同组 8 个头从「同一个 V 空间」取。但实证表明:同组内的 Q 差异足以产生足够不同的注意力权重,从同一个 V 里「加权」出足够不同的结果。

3.4 维度四:PCA 类比------8 个头 = 8 个主成分

把 K/V 看作对 token 信息的压缩编码

方案 每 token K/V 编码维度 压缩比
MHA 64×128=819264 \times 128 = 8192 64×128=8192 1×(无压缩)
GQA-8 8×128=10248 \times 128 = 1024 8×128=1024 1/8
MQA 1×128=1281 \times 128 = 128 1×128=128 1/64

类比 PCA(主成分分析):数据的主要变异方向往往只有少数几个。一个 token 的「可被检索属性」的主要变异方向大约就是 8 个左右:

yaml 复制代码
KV 头 0: 语法功能(主/谓/宾/定/状/补)
KV 头 1: 语义类别(人/物/动作/属性/关系)
KV 头 2: 位置/距离信息(句首/句中/句尾、远近)
KV 头 3: 指代/共指(代词指向谁)
KV 头 4: 主题/话题(属于哪个讨论主题)
KV 头 5: 搭配/共现(常与什么词一起出现)
KV 头 6: 情感/语气(正面/负面/中性)
KV 头 7: 信息结构(已知/新信息/焦点)

这 8 个维度已经张成了「一个 token 能向外界提供什么」的主要空间。再加更多头,边际信息增量趋近于零------就像 PCA 的第 9、第 10 个主成分解释方差极小。

3.5 维度五:实验证据(Ainslie et al., 2023 Ablation)

Google 在 GQA 原论文中的系统消融实验:

方案 KV 头数 质量(MMLU 等平均) KV Cache 大小 推理速度
MHA 64 基准(100%)
GQA-16 16 ≈基准(-0.1%) 1/4 ~3.5×
GQA-8 8 ≈基准(-0.1~0.3%) 1/8 ~5×
GQA-4 4 略降(-0.5~1%) 1/16 ~7×
MQA 1 明显降(-2~5%) 1/64 ~10×

结论 :8 个 KV 头是「质量几乎无损」与「显存/速度大幅改善」的最佳平衡点。从 64 减到 8,质量只掉 0.1~0.3%;但从 8 减到 1,质量下降明显(尤其 >30B 的大模型)。说明 8 个头确实捕获了 K/V 的几乎全部有效信息,剩下的 56 个头基本是冗余。

此外,Ainslie et al. 还验证了 uptraining 路径:已有 MHA 模型只需取各头 K/V 的均值初始化 GQA 的 KV 头,再微调几千步即可迁移,无需从头训练。LLaMA-2 70B 即以此方式发布。

3.6 维度六:甜蜜点------为什么是 8 而非其他

综合以上五个维度:

  • 太少(1~4):信息瓶颈太窄,大模型质量下降明显
  • 太多(16~64):KV Cache 节省不够,显存仍然是瓶颈
  • 8:恰好处于「质量无损」与「显存大幅改善」的帕累托最优点

这不是巧合,而是由自然语言的信息结构决定的------语言的「可检索属性」的主要维度大约就是 8 个量级(语法、语义、位置、指代、主题、搭配、情感、信息结构),更多的 KV 头只是对这些维度的冗余细分。

3.7 repeat_interleave 广播机制

GQA 在计算时,需要将 8 个 KV 头「逻辑展开」为 64 份,使每个 Q 头都能与对应的 KV 头做注意力:
Kexpanded=repeat_interleave(K,dim=head,repeats=8) K_{\text{expanded}} = \text{repeat\interleave}(K, \text{dim=head}, \text{repeats}=8) Kexpanded=repeat_interleave(K,dim=head,repeats=8)
Vexpanded=repeat_interleave(V,dim=head,repeats=8) V
{\text{expanded}} = \text{repeat\_interleave}(V, \text{dim=head}, \text{repeats}=8) Vexpanded=repeat_interleave(V,dim=head,repeats=8)

展开后的对应关系:
Kexpanded:,0:8,:=K:,0,:,Kexpanded:,8:16,:=K:,1,:,... K_{\text{expanded}}:, 0:8, : = K:, 0, :, \quad K_{\text{expanded}}:, 8:16, : = K:, 1, :, \quad \ldots Kexpanded:,0:8,:=K:,0,:,Kexpanded:,8:16,:=K:,1,:,...

实现上无需真的 expand------FlashAttention / PagedAttention 内核直接让同组 Q 头读取同一份 K/V 内存地址,零拷贝、零额外显存。逻辑展开只是为了概念清晰。

3.8 参数量对比

以 LLaMA-2 70B 单层为例:

矩阵 GQA-8 形状 GQA-8 参数量 MHA 形状 MHA 参数量
WQ W_Q WQ (8192,8192)(8192, 8192) (8192,8192) 67M (8192,8192)(8192, 8192) (8192,8192) 67M
WK W_K WK (8192,1024)(8192, 1024) (8192,1024) 8.4M (8192,8192)(8192, 8192) (8192,8192) 67M
WV W_V WV (8192,1024)(8192, 1024) (8192,1024) 8.4M (8192,8192)(8192, 8192) (8192,8192) 67M
WO W_O WO (8192,8192)(8192, 8192) (8192,8192) 67M (8192,8192)(8192, 8192) (8192,8192) 67M
合计 ~151M / 层 ~268M / 层

GQA-8 投影参数比 MHA 节省约 44%(151M vs 268M / 层)。80 层累计节省 ~9.4B 参数。


4. KV Cache 为什么需要

4.1 自回归重算的浪费

自回归生成中,每生成一个新 token,朴素做法都把整个序列重新跑一遍前向:

ini 复制代码
step t=3:  输入 [x₁,x₂,x₃]      → 计算 x₁,x₂,x₃ 的 K,V
step t=4:  输入 [x₁,x₂,x₃,x₄]   → 又算了一遍 x₁,x₂,x₃ 的 K,V !
step t=5:  输入 [x₁..x₅]         → x₁..x₄ 的 K,V 第三次被算 !

总计算量是 O(n3)O(n^3) O(n3) 级别的浪费(每步 O(n2)O(n^2) O(n2),共 nn n 步,且历史反复重算)。

4.2 因果掩码保证 K/V 不变 → 可缓存

因果掩码保证:token ii i 的 K、V 只依赖 x1...xi x_1 \ldots x_i x1...xi,与后续生成的 token 无关 。所以一旦算出, ki k_i ki、 vi v_i vi 在整个生成过程中永不改变 → 完全可以缓存复用。

注意:只缓存 K 和 V ,不缓存 Q。因为每一步只需要当前新 token 的 Q,旧 token 的 Q 用完即弃(它们的输出早已产出并被丢弃)。

这是 KV Cache 的理论合法性基础------没有因果掩码,历史 K/V 会因「看到」未来 token 而改变,缓存就失效了。


5. KV Cache 机制

5.1 append-only 缓存

arduino 复制代码
┌──────────────── KV Cache(每层各一份,随生成增长)──────────┐
│  K_cache: [k₁, k₂, ..., k_{t-1}]     ← append-only         │
│  V_cache: [v₁, v₂, ..., v_{t-1}]     ← append-only         │
└──────────────────────────────────────────────────────────┘

生成第 t 步:
  1. 仅对最新 token x_t 计算 q_t, k_t, v_t         (单 token,O(d²))
  2. append: K_cache ← [..., k_t],V_cache ← [..., v_t]
  3. 注意力: attn = softmax(q_t · K_cacheᵀ / √d_k) · V_cache
  4. 得到 x_t 的输出 → 过 FFN → 预测 x_{t+1}

5.2 复杂度: O(n3)→O(n2)O(n^3) \to O(n^2) O(n3)→O(n2)

  • 朴素自回归 :每步 O(t2⋅d)O(t^2 \cdot d) O(t2⋅d),共 nn n 步,总计 O(n3⋅d)O(n^3 \cdot d) O(n3⋅d)
  • KV Cache :每步 O(t⋅d)O(t \cdot d) O(t⋅d),共 nn n 步,总计 O(n2⋅d)O(n^2 \cdot d) O(n2⋅d)

长序列收益巨大: n=4096n=4096 n=4096 时,加速比约 4096 倍。

5.3 显存公式

KV_bytes=2×L×n×nkv_heads×dhead×precision_bytes\text{KV\bytes} = 2 \times L \times n \times n{\text{kv\heads}} \times d{\text{head}} \times \text{precision\_bytes} KV_bytes=2×L×n×nkv_heads×dhead×precision_bytes

符号 含义
22 2 K 和 V 各一份
LL L Transformer 层数
nn n 序列长度(prompt + 已生成)
nkv_heads n_{\text{kv\_heads}} nkv_heads KV 头数(MHA = hh h;GQA/MQA 更少)
dhead d_{\text{head}} dhead 每个头的维度
precision_bytes\text{precision\_bytes} precision_bytes FP16/BF16 = 2, FP8 = 1

5.4 实例计算

LLaMA-2 13B(MHA, FP16) L=40L=40 L=40, h=40h=40 h=40, dhead=128 d_{\text{head}}=128 dhead=128

  • 单 token KV: 2×40×40×128×2=819,200 B≈0.78 MB/token2 \times 40 \times 40 \times 128 \times 2 = 819{,}200 \text{ B} \approx 0.78 \text{ MB/token} 2×40×40×128×2=819,200 B≈0.78 MB/token
  • 2048 上下文: 2048×0.78≈1.6 GB2048 \times 0.78 \approx 1.6 \text{ GB} 2048×0.78≈1.6 GB(单请求!)
  • 并发 50: ∼80 GB\sim 80 \text{ GB} ∼80 GB------远超模型权重本身(13B × 2B ≈ 26 GB)

LLaMA-2 70B(GQA-8 vs MHA, FP16) L=80L=80 L=80, dhead=128 d_{\text{head}}=128 dhead=128, n=4096n=4096 n=4096

方案 KV 头数 每 token KV Cache 4096 token 上下文
MHA ( g=64g=64 g=64) 64 64×128×2×2=32 KB64 \times 128 \times 2 \times 2 = 32 \text{ KB} 64×128×2×2=32 KB 10.24 GB
GQA-8 ( g=8g=8 g=8) 8 8×128×2×2=4 KB8 \times 128 \times 2 \times 2 = 4 \text{ KB} 8×128×2×2=4 KB 1.28 GB
MQA ( g=1g=1 g=1) 1 1×128×2×2=0.5 KB1 \times 128 \times 2 \times 2 = 0.5 \text{ KB} 1×128×2×2=0.5 KB 0.16 GB

GQA-8 相比 MHA 把 KV Cache 缩小到 1/8,直接决定单卡能并发多少请求。

5.5 数据布局

典型实现中,每层维护两个张量(batch 版):

makefile 复制代码
K_cache: (batch, n_kv_heads, max_seq_len, d_head)
V_cache: (batch, n_kv_heads, max_seq_len, d_head)

预分配 max_seq_len 长度,用一个 cache_len 指针记录已填充到哪。

5.6 极简伪代码(单层、单请求)

python 复制代码
class KVCache:
    def __init__(self, max_len, n_kv_heads, d_head, dtype):
        self.K = zeros(n_kv_heads, max_len, d_head, dtype)  # 预分配显存
        self.V = zeros(n_kv_heads, max_len, d_head, dtype)
        self.len = 0                                        # 已缓存 token 数

    def append(self, k_new, v_new):                         # k_new: (n_kv_heads, 1, d_head)
        self.K[:, self.len, :] = k_new
        self.V[:, self.len, :] = v_new
        self.len += 1
        return self.K[:, :self.len, :], self.V[:, :self.len, :]

def attention_step(x_t, cache, W_q, W_k, W_v, W_o):
    q = x_t @ W_q                            # 仅当前 token 的 Q
    k = x_t @ W_k
    v = x_t @ W_v
    K, V = cache.append(k, v)                # 拿到含历史的完整 K/V
    scores = (q @ K.transpose(-1, -2)) / sqrt(d_head)   # (1, cache.len)
    attn = softmax(scores) @ V               # (1, d_head)
    return attn @ W_o

关键点:append 后返回的是从 0 到当前的全部 K/V 切片,Q 只有 1 行 → 因果掩码天然满足(当前 token 本就只能看它自己及之前)。


6. Prefill vs Decode

工程上,一次生成明确分为两个阶段,性能特征迥异:

6.1 Prefill 预填充阶段

sql 复制代码
┌──────────────── Prefill ──────────────────────────┐
│ 把用户输入的整个 prompt 一次性并行前向              │
│ • 目的: 建满 KV Cache + 产出第 1 个 token          │
│ • 特点: 计算密集(compute-bound),GPU 利用率高       │
│ • 操作: 矩阵 × 矩阵(GEMM),高度并行              │
│ • 指标: TTFT(Time To First Token,首 token 延迟)  │
└───────────────────────┬───────────────────────────┘
                        ▼

6.2 Decode 解码阶段

sql 复制代码
┌──────────────── Decode ──────────────────────────┐
│ 逐个 token 循环生成,每步 1 个 token               │
│ • 目的: 借助 KV Cache 一个个往外蹦                 │
│ • 特点: 访存密集(memory-bound),难并行              │
│ • 操作: 向量 × 矩阵(GEMV),算术强度低             │
│ • 指标: TPOT(Time Per Output Token)/ 吞吐量      │
└──────────────────────────────────────────────────┘

6.3 性能特征对比

维度 Prefill Decode
输入规模 整个 prompt(数百~数千 token) 1 个 token
计算类型 GEMM(矩阵×矩阵) GEMV(向量×矩阵)
瓶颈 Compute-bound(算力瓶颈) Memory-bound(显存带宽瓶颈)
GPU 利用率
对应指标 TTFT TPOT
优化方向 FlashAttention、算子融合 量化、Speculative Decoding

Decode 阶段的注意力退化为**「矩阵×向量」**:
S=Qnew⋅Kcache⊤:(h,1,dk)×(h,dk,t)→(h,1,t) S = Q_{\text{new}} \cdot K_{\text{cache}}^\top: \quad (h, 1, d_k) \times (h, d_k, t) \to (h, 1, t) S=Qnew⋅Kcache⊤:(h,1,dk)×(h,dk,t)→(h,1,t)

注意力矩阵从 (n,n)(n, n) (n,n) 退化为 (1,t)(1, t) (1,t)------一行向量,计算量 O(t⋅d)O(t \cdot d) O(t⋅d)。瓶颈变成读取 KV Cache 的显存带宽(memory-bound),而非计算。

实际体验:「首字慢(Prefill + TTFT),后续字匀速蹦(Decode)」。输入 prompt 越长,Prefill 越久,首字延迟越高。


7. 三大痛点

朴素的「预分配连续大张量」方式存在严重问题:

7.1 内部碎片(预留浪费)

不知道请求会生成多长,只能按 max_seq_len(如 2048)预分配。若实际只生成 100 个 token,其余 1948 个位置的显存全被占着却空闲

less 复制代码
传统连续分配(■=已用, □=预留浪费):
请求A: ■■■■□□□□□□□□□□□□   ← 生成4个,预留16个 → 75%浪费
请求B: ■■□□□□□□            ← 类似浪费

7.2 外部碎片

不同请求预留大小不同,请求结束后释放,显存被切成大小不一的空洞。新请求即使总空闲够用,也可能因没有足够大的连续块而分配失败。

css 复制代码
释放A后: ✕✕✕✕✕✕✕✕✕✕✕✕✕✕✕✕   ← 留下大空洞,难复用

7.3 无法共享

多个请求若有相同前缀(如同一 system prompt、并行采样 nn n 个候选),各自独立分配一份完全相同的 KV,无法共享,重复占用显存。

研究(vLLM 论文)测得:传统方式 KV 显存有效利用率常低至 20%~40%,60%~80% 被浪费。


8. PagedAttention

一句话本质 :PagedAttention 借用 OS 分页思想,把 KV Cache 从「必须连续」解放为「固定块 + 块表映射」,实现按需分配、任意复用、跨序列共享------显存利用率从 20%40% 跃升至 >90%,推理吞吐提升 24 倍。

§7 揭示了传统连续分配的三大痛点:内部碎片、外部碎片、无法共享。根源都在于一个假设------KV Cache 必须连续存放 。PagedAttention(Kwon et al., SOSP 2023)打破这个假设,做法与早期 OS 解决「程序必须用连续物理内存」的困境如出一辙------分页。对比如下:

操作系统概念 PagedAttention 对应
进程(Process) 一个请求/序列(Sequence)
虚拟内存页(Page) 逻辑块(Logical KV Block)
物理内存帧(Frame) 物理块(Physical KV Block)
页表(Page Table) 块表(Block Table)
按需分配页 按需分配 KV 块
写时复制(Copy-on-Write) KV 块的 CoW 共享

具体分三步:① 把显存切成等大物理块(block pool);② 每个序列用一张块表(Block Table)记录逻辑块→物理块的映射;③ 定制 CUDA 内核按块表逐块 gather K/V 计算注意力。

8.1 分块与块表

首先把显存切成等大物理块,再用块表管理映射。

物理块(Physical KV Block) :vLLM 默认 block_size = 16,整个显存被预切成 NN N 个等大的物理块,组成全局共享的「块池」:

yaml 复制代码
一个物理 KV 块的形状(每层):
  K_block: (block_size, n_kv_heads, d_head)
  V_block: (block_size, n_kv_heads, d_head)

块池:
┌────┬────┬────┬────┬────┬────┬────┬────┐
│ B0 │ B1 │ B2 │ B3 │ B4 │ B5 │ ...│ BN │   ← 全局共享
└────┴────┴────┴────┴────┴────┴────┴────┘

块大小权衡:块太小 → 块表大、寻址频繁;块太大 → 内部碎片回升。默认 16 是折中。

块表(Block Table):每个序列拥有一张 Block Table,记录逻辑块到物理块的映射。逻辑上连续,物理上分散------这就是 PagedAttention 的精髓:

ini 复制代码
序列 A 的 tokens:  [t0 t1 t2 t3 | t4 t5 t6 t7 | t8 t9]     (block_size=4)
                   逻辑块0        逻辑块1        逻辑块2(未满)

Block Table (A):
  逻辑块 0 → 物理块 B3
  逻辑块 1 → 物理块 B7
  逻辑块 2 → 物理块 B1   (只用了 2/4 槽位)

物理显存实际布局(不连续!):
  B1: [t8 t9 _ _]     B3: [t0 t1 t2 t3]     B7: [t4 t5 t6 t7]

有了这个机制,§7 的三大痛点迎刃而解:按需分配使内部碎片 <4%(仅浪费最后一块未满部分);所有物理块等大,外部碎片 = 0;多序列的块表可指向同一物理块,支持前缀/采样共享。

8.2 块管理器:CoW 共享与调度

分块与块表解决了「怎么放」,但多序列场景还需要解决「怎么共享」和「怎么管」。vLLM 用 Block Manager 统一管理这两件事。

前缀共享与写时复制(CoW):多个请求共享相同 system prompt 时,前缀 token 的 KV 完全相同------只存一份物理块,多个序列的块表都指向它:

less 复制代码
system prompt "你是一个助手..." → 逻辑块0,1
请求A 块表: [B5, B8, B2, ...]
请求B 块表: [B5, B8, B9, ...]   ← 前两块与 A 共享同一物理块!

共享块是只读共享的。当某序列需要修改一个共享块时,才触发复制(Copy-on-Write):

less 复制代码
初始: 序列A、B 共享物理块 B5(引用计数 ref=2)

序列A 要在 B5 所在的逻辑块追加新 token:
  1. 分配新物理块 B12
  2. 把 B5 内容复制到 B12
  3. A 的块表改指 B12,在 B12 写入新 token
  4. B5 的 ref 减为 1,B 继续独享 B5

→ 只在"真正要写且被共享"时才复制,避免不必要的拷贝

引用计数(ref count)管理每个物理块的引用,ref=0 时回收进空闲块池。这与 OS 的 CoW fork、页面回收如出一辙。

Block Manager 与调度:所有上述操作由 Block Manager 统一调度:

scss 复制代码
┌──────────────── Block Manager ────────────────┐
│  free_list: 空闲物理块队列                       │
│  ref_count[phys]: 每个物理块的引用计数            │
│  block_tables[seq_id]: 每个序列的逻辑→物理映射    │
│                                                 │
│  allocate(seq): 从 free_list 取块,建/扩块表      │
│  append_token(seq): 当前块满则申请新块           │
│  fork(seqA→seqB): 复制块表,共享块 ref++(CoW)   │
│  free(seq): 块表所有块 ref--,归零则回 free_list  │
└─────────────────────────────────────────────────┘

当显存不足以容纳所有活跃请求时,调度器还支持抢占与换出Swap (把 KV 块换出到 CPU 内存,需要时换回)和 Recompute (丢弃 KV,重新 Prefill,用算力换显存)。配合连续批处理(Continuous Batching)------请求随时加入/离开批次,不必等整批结束,GPU 始终填满,吞吐大幅提升。

8.3 注意力内核

以上解决了「怎么存」,最后一个问题:K/V 现在分散在非连续物理块中,标准注意力内核假设连续存放,怎么办?PagedAttention 定制了 CUDA 内核,按块计算:

python 复制代码
for 逻辑块 j in 序列的所有块:
    phys = block_table[j]              # 查块表得物理块号
    K_j = physical_K[phys]             # 取出该物理块的 K
    V_j = physical_V[phys]
    scores_j = q_t · K_jᵀ / √d_head    # 只和这一块算
    # 累积到全局 softmax(用 online-softmax 数值稳定地增量归一化)
    partial = softmax_accumulate(scores_j)
    out += partial · V_j
return out

三个子技术环环相扣:

首先,按块取数(Block-wise Gather)------内核必须解决「怎么读」的问题。标准注意力内核假设 K/V 在显存中连续存放,可以顺序读取(coalesced access),这是 GPU 最高效的访存模式。但物理块不连续,内核必须"跳跃式"读取。

给定一个 token 的逻辑位置 pospos pos,物理地址只需一次整数除法和取模即可算出:

ini 复制代码
block_idx = pos // block_size        # 逻辑块号
offset    = pos %  block_size        # 块内偏移
phys_addr = physical_blocks[ block_table[block_idx] ][offset]

在 GPU 上,每个 CUDA thread block 负责处理一个 query token:线程协作从 Block Table 查出物理块号,再 gather 对应的 K/V 向量。

ini 复制代码
Query token q_t 需要读取 3 个逻辑块的 K/V(block_size=4):

Block Table:  [B3, B7, B1]      ← 逻辑块 → 物理块

物理显存布局(不连续!):
  ┌─────────┐    ┌─────────┐    ┌─────────┐
  │  B3      │    │  B7      │    │  B1      │
  │ K0 K1    │    │ K4 K5    │    │ K8 K9    │
  │ K2 K3    │    │ K6 K7    │    │ -- --    │  (最后一块未满)
  └─────────┘    └─────────┘    └─────────┘
       ↑               ↑               ↑
  地址 0x1A00      地址 0x3F00      地址 0x0C00
       │               │               │
       └───────────────┴───────────────┘
                  gather 到 SRAM
                      ↓
              q_t · [K0..K9]ᵀ / √d_k    ← 在 SRAM 中连续计算

block_size=16block\_size=16 block_size=16 保证了每个块内的读取仍是连续的(粒度可控),不会退化为完全随机访问。代价仅是 Block Table 的一次查表(整数索引),换来的是显存管理的完全灵活性

取数方式变了,softmax 也必须跟着变。

Online Softmax :标准 softmax 需要两遍遍历 :先求全局 max⁡\max max,再求 ∑\sum ∑ 并归一化。这意味着必须先将所有 score 存下来才能开始归一化------对分块计算不友好。Online Softmax 维护两个累积量------ mm m(当前最大值)和 dd d(当前指数和),每处理一块就增量更新,无需存储完整的 score 矩阵:

ini 复制代码
# 初始化
m = -∞
d = 0
out = 0

# 处理第 j 块
scores_j = q · K_jᵀ / √d_k                          # 与第 j 块的点积 (block_size,)
m_new = max(m, max(scores_j))                        # 更新最大值
d_new = d · exp(m - m_new) + Σ exp(scores_j - m_new) # 更新指数和(校正历史)
out   = out · (d · exp(m - m_new) / d_new)
        + exp(scores_j - m_new) / d_new · V_j         # 更新输出
m = m_new
d = d_new

关键校正步骤 :当新块的 score 比当前 mm m 更大时,历史累积的 dd d 和 outout out 已经按旧的 mold m_{old} mold 计算过指数。此时必须乘以 exp⁡( mold − mnew ) \exp(m_{old} - m_{new}) exp(mold−mnew) 进行"回溯校正"------把之前"高估"的权重缩回来。

与 FlashAttention 的关系 :PagedAttention 复用的是 FlashAttention 中的 Online Softmax 子技术(增量维护 mm m 与 dd d,边算边归一化),而非整个 FlashAttention 算法。两者正交------FlashAttention 优化单次注意力的 IO(通过 tiling 减少 HBM 读写),PagedAttention 优化 KV Cache 的显存管理(消除碎片),可叠加使用。

前两步让分块计算可行,但性能还差一步。

融合(Kernel Fusion):GPU 计算的瓶颈往往不是算力(compute-bound),而是数据搬运(memory-bound)。如果将注意力拆成多个独立 kernel,中间结果会反复在 HBM 和 SRAM 之间搬运:

csharp 复制代码
              非融合                              融合
        ┌───────────────┐               ┌───────────────┐
        │     HBM       │               │     HBM       │
        │  K/V scores   │               │   output      │
        │  weights out  │               └───────┬───────┘
        └──┬──┬──┬──┬───┘                       │ 1次往返
           │  │  │  │ 4次往返            ┌───────▼───────┐
        ┌──▼──▼──▼──▼───┐               │    SRAM       │
        │    SRAM        │               │ gather→QKᵀ   │
        │ (每步只算一个)  │               │ →softmax→Σ   │
        └───────────────┘               └───────────────┘

融合将 gather + 点积 + Online Softmax + 加权求和全部放在单个 kernel 内,中间结果留在 SRAM/寄存器中,将 4 次 HBM 往返压缩为 1 次。

SRAM 预算 :A100 每个 SM 有 192KB SRAM,足以容纳一个 query token 对所有 block 的部分累积结果------ mm m、 dd d、 outout out 各为 dhead d_{head} dhead 维向量(如 dhead =128 d_{head}=128 dhead=128 时仅占 3×128×4B=1.5KB3 \times 128 \times 4\text{B} = 1.5\text{KB} 3×128×4B=1.5KB),加上当前块的 K/V( 16×128×4B×2=16KB16 \times 128 \times 4\text{B} \times 2 = 16\text{KB} 16×128×4B×2=16KB),远在 SRAM 容量之内。这也是 PagedAttention 内核能在引入 Block Table 寻址开销的同时仍保持高性能的关键原因。

至此,PagedAttention 的完整链路闭合:分块存储解决碎片、块表提供灵活映射、CoW 实现跨序列共享、融合内核保证计算性能------代价仅是一点点寻址开销。

8.4 效果

以上四个环节------分块、块表、CoW、融合内核------协同工作,最终效果:

指标 传统方式 PagedAttention
显存利用率 20%~40% >90%
吞吐提升 基准 2~4 倍(vLLM 论文)
内部碎片 60%~80% <4%
外部碎片 严重 0
前缀共享 不支持 支持(CoW)

9. 前沿进展

9.1 KV Cache 量化(FP8 / INT4)

将 KV Cache 从 FP16/BF16(2 字节)量化为更低精度:

精度 字节数 显存节省 质量影响
FP16/BF16 2 基准 无损
FP8 (E4M3) 1 50% 极小(<0.5% MMLU 下降)
INT4 0.5 75% 中等(1~3% 下降)

KV Cache 量化与 GQA 正交叠加 :GQA-8 已将 KV 头数压缩到 1/8,再叠加 FP8 量化,KV Cache 总开销降至 MHA FP16 的 1/16。这对于超长上下文(128K+ tokens)尤为重要。

实现挑战:K/V 的分布在不同层、不同头之间差异较大,需要 per-channel 或 per-token 的量化粒度以保持精度。

9.2 Prefix Caching(前缀缓存)

当多个请求共享相同前缀(system prompt、few-shot examples)时,跨请求复用已计算好的 KV 块,避免重复 Prefill:

  • RadixAttention(SGLang):用前缀树(Radix Tree)自动识别共享前缀,匹配到的前缀直接引用已有 KV 块
  • Automatic Prefix Caching(vLLM):基于 block hash 匹配,相同内容的 block 自动共享
  • 效果:对于共享 1024 token system prompt 的场景,Prefill 时间可减少 60%~80%

与 PagedAttention 的 CoW 机制天然配合:共享前缀的物理块被多序列引用(ref > 1),只有分叉后的新 token 才需要独立分配块。

9.3 Speculative Decoding 与 KV Cache 的交互

Speculative Decoding(投机解码)用一个小模型(draft model)快速生成 kk k 个候选 token,再用大模型一次性验证:

yaml 复制代码
Draft Model: 快速生成 [t₁, t₂, ..., tₖ]
Target Model: 一次性验证全部 k 个 token(Prefill 模式)
  → 接受正确的前 m 个,拒绝剩余,从第 m+1 个重新 draft

与 KV Cache 的交互

  • Draft 阶段:小模型维护自己的 KV Cache,快速 append
  • Verify 阶段:大模型对 kk k 个 token 做批量 Prefill,无需逐步 append KV Cache
  • 被拒绝的 token:需要回滚大模型 KV Cache(丢弃被拒 token 对应的 K/V 条目)
  • 被接受的 token:KV Cache 保留,继续下一轮 draft

关键优化 :验证阶段利用 batch 计算,将 Decode(memory-bound)转化为 Prefill(compute-bound),大幅提升 GPU 利用率。典型加速比:2~3 倍 TPOT 改善,且输出分布与朴素解码严格一致(无损加速)。

9.4 Cross-Layer Attention(CLA)

传统注意力是层内 的------每层的 Q 只能看到同层的 K/V。Cross-Layer Attention 允许跨层共享 K/V

  • 思路 :第 ll l 层的 Q 可以访问第 l−1l-1 l−1 层(或更早层)的 K/V Cache,无需在本层重新计算
  • 动机:相邻层的 K/V 表示高度相似(已有实证支持),跨层复用可进一步压缩 KV Cache
  • 代表工作:CLA(Cross-Layer Attention)、MLA(Multi-head Latent Attention, DeepSeek-V2)
  • MLA 特别值得注意:将 KV 压缩为低秩潜向量(latent),推理时只缓存潜向量,按需解压缩出 K/V,KV Cache 可进一步缩小到 GQA 的 5%~13%

9.5 vLLM 性能基准数据

基于 vLLM(PagedAttention 引擎)的典型性能表现(A100 80GB, FP16):

模型 并发数 吞吐量(tokens/s) TTFT (ms) TPOT (ms)
LLaMA-2 7B 64 ~3,500 ~80 ~18
LLaMA-2 13B 32 ~2,000 ~120 ~25
LLaMA-2 70B(GQA-8, 4×A100) 16 ~800 ~300 ~45
LLaMA-3 8B 64 ~4,000 ~70 ~15

:以上为典型参考值,实际性能受 prompt 长度、生成长度、批处理策略、GPU 型号等因素影响。

PagedAttention 的核心贡献 :在相同硬件上,相比 HuggingFace Transformers 早期方案,vLLM 的吞吐量提升 2~4 倍,主要归功于显存利用率从 20%~40% 提升至 >90%,从而容纳更大 batch。


总结

主题 核心要点
GQA 多个 Q 头共享少量 KV 头,KV Cache 缩至 1/8,质量几乎无损(<0.3%),已成为事实标准
KV Cache 缓存历史 K/V 避免重复计算, O(n3)→O(n2)O(n^3) \to O(n^2) O(n3)→O(n2),是推理显存的主要消耗
Prefill/Decode Prefill = compute-bound(建 Cache),Decode = memory-bound(逐 token 生成)
三大痛点 内部碎片、外部碎片、无法共享------传统连续分配的致命缺陷
PagedAttention OS 分页思想 → 分块 + 块表 + CoW → 利用率 >90%,吞吐提升 2~4×
前沿 KV 量化、Prefix Caching、Speculative Decoding、CLA/MLA 持续压缩 KV Cache 开销

一句话总结

GQA 从模型结构 层面将 KV 冗余压缩到 1/8,KV Cache 从算法 层面消除重复计算,PagedAttention 从系统层面将显存利用推向极致------三者层层递进,共同解决了大模型推理中「显存」这一核心瓶颈,使得在有限硬件上服务海量并发请求成为可能。

相关推荐
DogDaoDao1 小时前
OpenBrowser 深度解析:让 AI 真正「用上」浏览器的自主代理框架
人工智能·程序员·大模型·github·web·ai工具·openbrowser
_Jimmy_1 小时前
Tool Calling 与 Function Calling 区别
人工智能·python·langchain
ITmaster07311 小时前
告别 IDE?Android CLI 来了,开发进入 AI Agent 时代
android·ide·人工智能
陈明勇1 小时前
一篇文章,多种表达:我用 Seed Evolving 生成知识卡片
人工智能
墨舟的AI笔记1 小时前
ECS 中的确定性随机与回放:让帧同步在 DOTS 上成立
人工智能
图特摩斯科技2 小时前
本体智能应用案例实践分享:汽车零部件库存优化与召回应急保障
人工智能·汽车·palantir·ontology·ontoflow·ontoos
专业工业电源打工人2 小时前
F0505S-2WR3 适配优选 钡特电源 DF2-05S05LS|2W 隔离 DC-DC 模块电源5V转5V硬件选型参数规格解析
大数据·网络·人工智能
a1117762 小时前
三色软糖坠落玻璃池 THreeJS kimi
前端·人工智能·threejs
xian_wwq2 小时前
【学习笔记】解剖 Claude Code —— Anthropic 的 Harness 参考实现-09/15
人工智能·笔记·学习