从 MHA 到 MQA 和 GQA:KV Cache 为什么会变小

看到 MQA 的“共享 K/V”时,一个直接的问题是:标准 MHA 的多个 Head 明明接收同一批 Token,为什么还要专门强调共享?

这里混在一起的是两种共享。MHA 的多个 Head 共享同一个输入隐藏状态 X,但各自生成独立的 Q、K、V。MQA 进一步让所有 Query Head 使用同一组投影后的 K/V。GQA 位于两者之间,每组 Query Head 共享一组 K/V。

从一次序列预测开始

为避开具体 tokenizer 的切分差异,把下面每一项视为一个抽象 Token:

上下文 Token:我 | 爱 | 吃
预测目标 Token:苹果

在某个 Transformer 层中,三个输入 Token 的隐藏状态组成矩阵 X。若序列长度为 T,模型隐藏维度为 d_model,则:

X 的形状:[T, d_model]

自注意力层用三组可学习参数投影出 Q、K、V:

$$ Q=XW_Q,\qquad K=XW_K,\qquad V=XW_V $$

因此,每个 Token 在每一层中都有一个 Q 向量、一个 K 向量和一个 V 向量。多头注意力会把这些向量继续按 Head 拆分。

假设:

d_model = 8
Query Head 数 H = 4
每个 Head 的维度 d_head = 2

单个 Token 的 Q 可以写成:

q = [q_1, q_2, q_3, q_4]

每个 q_h 都是二维向量
整个 q 仍是八维向量

对完整序列,MHA 中 Q、K、V 的形状均为:

[T, 4, 2]

所以“每个 Token 的 QKV 是向量”和“每个 Token 有多组 QKV Head”可以同时成立。前一种说法观察完整向量,后一种说法观察拆分后的子向量。

MHA 怎样计算当前 Token

标准 MHA 中,每个 Head 都有自己的投影参数。以第 h 个 Head 为例:

$$ q_t^{(h)}=x_tW_Q^{(h)},\qquad k_t^{(h)}=x_tW_K^{(h)},\qquad v_t^{(h)}=x_tW_V^{(h)} $$

对于上下文中的“吃”,第 h 个 Query Head 会和同一个 Head 下的三个 Key 做点积:

$$ s_j^{(h)}= \frac{q_{\text{吃}}^{(h)}{k_j^{(h)}}^T}{\sqrt{d_{head}}}, \qquad j\in\{\text{我},\text{爱},\text{吃}\} $$

经过 softmax 得到注意力权重,再对该 Head 的 Value 加权求和:

$$ \alpha_j^{(h)}= \operatorname{softmax}\left(s_j^{(h)}\right), \qquad o_{\text{吃}}^{(h)}= \sum_j\alpha_j^{(h)}v_j^{(h)} $$

下面用二维教学数值展示一个 Head 的计算。这些数值只用于展开公式:

q_吃 = [1, 1]

k_我 = [1, 0]
k_爱 = [0, 1]
k_吃 = [1, 1]

v_我 = [1, 0]
v_爱 = [0, 1]
v_吃 = [1, 1]

缩放点积为:

score_我 = 0.707
score_爱 = 0.707
score_吃 = 1.414

softmax 后的权重约为:

[0.248, 0.248, 0.503]

该 Head 的输出为:

0.248 × [1, 0] + 0.248 × [0, 1] + 0.503 × [1, 1]
= [0.752, 0.752]

四个 Head 独立执行同样的过程。随后把四个输出拼接,并乘输出投影矩阵:

$$ o_t= \operatorname{Concat}\left( o_t^{(1)},o_t^{(2)},o_t^{(3)},o_t^{(4)} \right)W_O $$

因果掩码限制“吃”只能读取“我”“爱”“吃”的 K/V。这个结果再经过残差连接、FFN 和后续 Transformer 层。最后一层在“吃”位置产生隐藏状态 h_吃,LM Head 将它投影到词表,得到下一个 Token 的概率:

$$ P(\text{next token})= \operatorname{softmax}\left(h_{\text{吃}}W_{vocab}\right) $$

“苹果”由完整模型的最终概率分布产生,前面的二维计算只展示单层、单 Head 的信息聚合过程。

MHA 共享了什么

MHA 中有两类复用:

  1. 所有 Head 接收同一个输入隐藏状态 X。
  2. 同一套投影参数会应用于序列中的每个 Token。

Head 之间仍有独立的参数块:

第 1 个 Head:W_Q1、W_K1、W_V1
第 2 个 Head:W_Q2、W_K2、W_V2
第 3 个 Head:W_Q3、W_K3、W_V3
第 4 个 Head:W_Q4、W_K4、W_V4

工程实现通常把四个小矩阵拼成一个大矩阵,一次完成投影:

W_K = [W_K1 | W_K2 | W_K3 | W_K4]
K = X W_K

随后把 K 从 [T, 8] reshape 为 [T, 4, 2]。一个大 K 张量只是融合计算和存储形式,拆开后仍是四组独立的 K Head。V 和 Q 同理。

MHA 的 Head 对应关系为:

Q1 使用 K1、V1
Q2 使用 K2、V2
Q3 使用 K3、V3
Q4 使用 K4、V4

因此,MHA 共享输入 X,同时保留每个 Head 独立的投影结果。

KV Cache 为什么只保存 K 和 V

Prefill 处理完“我 | 爱 | 吃”后,每一层已经得到了这些 Token 的 K/V。生成下一个 Token 时,当前 Query 需要再次读取全部历史 K/V:

当前 q 读取:k_我、k_爱、k_吃
当前 q 加权:v_我、v_爱、v_吃

历史 Query 已经完成各自位置的注意力计算,后续步骤不会再次读取它们。推理引擎因此保存 K/V,并丢弃历史 Q。

对 MHA 而言,每个 Token、每一层都需要保存 H 组 K 和 H 组 V。忽略张量排列差异,缓存元素数量为:

$$ N_{KV}= B\times L\times T\times H_{KV}\times d_{head}\times 2 $$

其中 B 是 Batch Size,L 是层数,T 是缓存 Token 数,最后的 2 代表 K 和 V。MHA 中 H_KV 等于 Query Head 数 H。

MQA 如何共享 K 和 V

MQA 保留多个 Query Head,但每层只生成一组 K 和一组 V:

$$ q_t^{(h)}=x_tW_Q^{(h)},\qquad k_t=x_tW_K,\qquad v_t=x_tW_V $$

对于 4 个 Query Head:

Q 的形状:[T, 4, 2]
K 的形状:[T, 1, 2]
V 的形状:[T, 1, 2]

对应关系变为:

Q1 使用 K、V
Q2 使用 K、V
Q3 使用 K、V
Q4 使用 K、V

K 和 V 仍由可学习矩阵完成投影。W_K 与 W_V 是两套参数,各自产生一个 d_head 维向量。区别在于,这组投影结果由全部 Query Head 使用。

注意力 Kernel 可以广播这组 K/V,也可以在计算视图中临时展开。KV Cache 仍只保存一组,临时展开不会把缓存永久复制成四组。

GQA 如何分组共享

GQA 把 H 个 Query Head 分成 G 组,每组对应一个 K Head 和一个 V Head。G 也就是 KV Head 数。

假设有 8 个 Query Head、2 个 KV Head:

Q1、Q2、Q3、Q4 使用 K1、V1
Q5、Q6、Q7、Q8 使用 K2、V2

张量形状为:

Q 的形状:[T, 8, d_head]
K 的形状:[T, 2, d_head]
V 的形状:[T, 2, d_head]

每个 KV Head 服务 H / G 个 Query Head。三种注意力可以用同一组参数描述:

结构 Query Head 数 KV Head 数 每组 KV 服务的 Query Head 数
MHA H H 1
GQA H G H / G
MQA H 1 H

当 G=H 时得到 MHA;当 1<G<H 时是 GQA;当 G=1 时得到 MQA。

三种结构在同一次 Decode 中的区别

假设上下文已经包含三个 Token,模型现在处理第四个 Token。每层都会为第四个 Token 生成多个 Query:

q4_1、q4_2、q4_3、q4_4

三种结构读取历史缓存的方式分别是:

MHA:
q4_1 读取历史 K1、V1
q4_2 读取历史 K2、V2
q4_3 读取历史 K3、V3
q4_4 读取历史 K4、V4

GQA,两个 KV Head:
q4_1、q4_2 读取历史 K1、V1
q4_3、q4_4 读取历史 K2、V2

MQA:
q4_1、q4_2、q4_3、q4_4 读取同一份历史 K、V

三种结构都保留多组 Query,也都输出多组 Head 结果。变化集中在 K/V 的投影数量、缓存数量和读取带宽。

Llama 2 70B 的缓存账

Meta 的 Llama 2 70B 配置包含 80 层、64 个 Query Head、8 个 KV Head,隐藏维度为 8192,因此 d_head=128。它采用 GQA。

假设 Batch Size 为 1、上下文长度为 4096,KV Cache 使用 FP16 或 BF16,每个元素占 2 字节:

$$ \text{KV Cache Bytes}= 2\times L\times T\times H_{KV}\times d_{head}\times 2 $$

第一个 2 表示 K 和 V,最后一个 2 表示每个 16 位元素占 2 字节。

结构 KV Head 数 计算结果
MHA,同尺寸假设 64 10 GiB
GQA,Llama 2 70B 配置 8 1.25 GiB
MQA,同尺寸假设 1 160 MiB

GQA 的缓存是同尺寸 MHA 的八分之一;MQA 是六十四分之一。Batch Size 和上下文长度增加时,KV Cache 按比例增长。

容易混淆的四件事

第一,QKV 没有一份先生成的“原始版本”等待各个 Head 再加工。自注意力从当前层输入 X 出发,通过 Head 对应的参数块直接得到 Q、K、V。

第二,MHA 的融合投影矩阵只是一种实现方式。大矩阵内部的不同列块对应不同 Head 参数,因此每个 Head 仍拥有自己的 K/V。

第三,MQA 仍会投影 K 和 V。它只把 K Head 和 V Head 的数量降为 1,同时保留多个 Query Head。

第四,GQA 的 G 指 KV Head 数。它把 MHA 和 MQA 放在同一条连续轴上,用 KV Head 数控制缓存开销与 Head 间共享程度。

参考资料