看到 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:
因此,每个 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 为例:
对于上下文中的“吃”,第 h 个 Query Head 会和同一个 Head 下的三个 Key 做点积:
经过 softmax 得到注意力权重,再对该 Head 的 Value 加权求和:
下面用二维教学数值展示一个 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 独立执行同样的过程。随后把四个输出拼接,并乘输出投影矩阵:
因果掩码限制“吃”只能读取“我”“爱”“吃”的 K/V。这个结果再经过残差连接、FFN 和后续 Transformer 层。最后一层在“吃”位置产生隐藏状态 h_吃,LM Head 将它投影到词表,得到下一个 Token 的概率:
“苹果”由完整模型的最终概率分布产生,前面的二维计算只展示单层、单 Head 的信息聚合过程。
MHA 共享了什么
MHA 中有两类复用:
- 所有 Head 接收同一个输入隐藏状态
X。 - 同一套投影参数会应用于序列中的每个 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。忽略张量排列差异,缓存元素数量为:
其中 B 是 Batch Size,L 是层数,T 是缓存 Token 数,最后的 2 代表 K 和 V。MHA 中 H_KV 等于 Query Head 数 H。
MQA 如何共享 K 和 V
MQA 保留多个 Query Head,但每层只生成一组 K 和一组 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 字节:
第一个 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 间共享程度。
参考资料
- Attention Is All You Need,MHA 与缩放点积注意力的原始定义。
- Fast Transformer Decoding: One Write-Head is All You Need,MQA 的结构与解码带宽动机。
- GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints,GQA 的分组方式与实验结果。
- Meta Llama 模型配置,Llama 2 70B 的层数、隐藏维度和 Head 配置。
- Meta Llama 2 Model Card,Llama 2 70B 使用 GQA 的说明。