MQA 多查询注意力(Multi-Query Attention)
概述
传统多头注意力(MHA)为每个头分配独立的 Q、K、V 矩阵。
MQA 的核心改动是:所有头共享同一组 K、V 矩阵,但每个头保留自己独立的 Q 矩阵。
- 好处:K、V 只需计算和存储一次,大幅减少 KV Cache 显存和内存访问。
- 代价:表达能力略有下降,但推理速度显著提升。
数值例子
设
text
d_model = 4
头数 h = 2
每个头 d_k = d_v = 2
输入序列一个 token:x = [1, 0, 1, 0]
MHA 做法:
每个头有自己的 WQ、WK、WVW_Q、W_K、W_VWQ、WK、WV,分别得到 Q1,K1,V1Q1,K1,V1Q1,K1,V1 和 Q2,K2,V2。
MQA 做法:
所有头共享一组 K、V,但 Q 仍然独立。
假设共享的 K、V 权重为单位阵,Q 权重也为单位阵(简化)。
- 共享的 K=x=1,0,1,0K = x = 1, 0, 1, 0K=x=1,0,1,0(切成两个头?实际上 MQA 中 K、V 的维度通常等于每个头的维度,但所有头共享同一组,所以 K 向量维度为 dk=2d_k=2dk=2?这里需要明确:MQA 中 K、V 的投影矩阵是共享的,输出维度为 dkd_kdk。我们简单设 dk=2d_k=2dk=2,那么共享的 K=xWK=1,0K = xW_{K} = 1,0K=xWK=1,0(取前两维),V = 1,0。
- 头 1 的 Q1 = xWQ1xW_{Q1}xWQ1 = 1,0
- 头 2 的 Q2 = xWQ2xW_{Q2}xWQ2 = 1,0(假设相同,便于计算)
计算头 1 的注意力:
- Q1⋅K=1,0⋅1,0=1Q1·K = 1,0·1,0 = 1Q1⋅K=1,0⋅1,0=1
- 缩放:1/√2≈0.70711/√2 ≈ 0.70711/√2≈0.7071
- softmax 只有一个 key,权重为 1
- 输出 Z1=V=1,0Z1 = V = 1,0Z1=V=1,0
头 2 同理,Z2=1,0Z2 = 1,0Z2=1,0。
拼接后得到 1,0,1,01,0,1,01,0,1,0。
对比 MHA: MHA 需要为每个头计算独立的 K、V,共 2 组;MQA 只计算 1 组 K、V,节省了一半 KV 存储和计算。
GQA 分组查询注意力(Grouped Query Attention)
概述
GQA 是 MHA 和 MQA 的折中:
- 将 Q 的头分成 GG 个组;
- 每组共享一组 K、V;
- 每个 Q 头仍然独立。
当 G=h 时退化为 MHA;当 G=1 时退化为 MQA。
数值例子
设:
text
头数 h = 4
分组数 G = 2
每组 2 个 Q 头共享一组 K、V
d_k = d_v = 2
输入 x = [1, 0, 1, 0]
- 组 1:Q1,Q2共享K1,V1Q1, Q2 共享 K1, V1Q1,Q2共享K1,V1
- 组 2:Q3,Q4共享K2,V2Q3, Q4 共享 K2, V2Q3,Q4共享K2,V2
简化设所有投影为单位阵,则:
- K1 = V1 = 1,0
- K2 = V2 = 0,1(假设不同组取不同部分)
计算组 1 中头 1:
- Q1 = 1,0
- Q1·K1 = 1
- 缩放后 0.7071,softmax 权重 1
- Z1 = V1 = 1,0
组 1 中头 2:
- Q2 = 1,0
- 同样输出 Z2 = 1,0
组 2 中头 3:
- Q3 = 0,1
- Q3·K2 = 1
- Z3 = V2 = 0,1
头 4 同理 Z4 = 0,1
最终拼接:1,0,1,0,0,1,0,1。
KV Cache 对比:
- MHA:4 组 K、V
- GQA:2 组 K、V
- MQA:1 组 K、V
GQA 在表达能力和效率之间取得平衡。
为什么 GQA 有效?------ KV Cache 算术强度视角
算术强度定义
算术强度 = 计算量 (FLOPs) / 内存访问量 (Bytes)。
自回归推理时,每生成一个 token:
- 计算量:主要来自 Q 与整个 KV Cache 的注意力计算,约为 O(L⋅d)O(L \cdot d)O(L⋅d);
- 内存访问:需要读取整个 KV Cache,大小约为 O(L⋅d⋅num_kv_heads)O(L \cdot d \cdot \text{num\_kv\_heads})O(L⋅d⋅num_kv_heads)。
因此算术强度大致为:
算术强度∝L⋅dL⋅d⋅num_kv_heads=1num_kv_heads \text{算术强度} \propto \frac{L \cdot d}{L \cdot d \cdot \text{num\_kv\_heads}} = \frac{1}{\text{num\_kv\_heads}} 算术强度∝L⋅d⋅num_kv_headsL⋅d=num_kv_heads1
- MHA: num_kv_heads = h
- MQA: num_kv_heads = 1
- GQA: num_kv_heads = G
减少 KV 头数,分母变小,算术强度提高,GPU 等待内存搬运的时间减少,吞吐量提升。
数值示例
设:
text
序列长度 L = 1024
模型维度 d = 4096
头数 h = 32
每个头维度 d_k = 128
数据类型 fp16(2 字节)
MHA 的 KV Cache 大小:
2×L×h×dk×2=2×1024×32×128×2=16,777,216 bytes≈16 MB 2 \times L \times h \times d_k \times 2 = 2 \times 1024 \times 32 \times 128 \times 2 = 16,777,216 \text{ bytes} \approx 16 \text{ MB} 2×L×h×dk×2=2×1024×32×128×2=16,777,216 bytes≈16 MB
GQA (G=8) 的 KV Cache 大小:
2×1024×8×128×2=4,194,304 bytes≈4 MB 2 \times 1024 \times 8 \times 128 \times 2 = 4,194,304 \text{ bytes} \approx 4 \text{ MB} 2×1024×8×128×2=4,194,304 bytes≈4 MB
MQA 的 KV Cache 大小:
2×1024×1×128×2=524,288 bytes≈0.5 MB 2 \times 1024 \times 1 \times 128 \times 2 = 524,288 \text{ bytes} \approx 0.5 \text{ MB} 2×1024×1×128×2=524,288 bytes≈0.5 MB
GQA 将 KV Cache 减少到 MHA 的 1/4,算术强度提高 4 倍,推理速度显著提升,且性能损失很小。
稀疏 / 滑动窗口注意力
概述
传统自注意力每个 token 关注所有 token,计算量 O(L2)O(L^2)O(L2)。
稀疏/滑动窗口注意力让每个 token 只关注局部窗口内的 token,例如只看前后 www 个 token,计算量降为 O(L⋅w)O(L \cdot w)O(L⋅w)。
数值例子
设序列长度 L=6L = 6L=6,窗口大小 w=3w = 3w=3(每个 token 只看自己和前后各 1 个)。
注意力矩阵(1 表示关注,0 表示不关注):
text
token: 0 1 2 3 4 5
0: 1 1 0 0 0 0
1: 1 1 1 0 0 0
2: 0 1 1 1 0 0
3: 0 0 1 1 1 0
4: 0 0 0 1 1 1
5: 0 0 0 0 1 1
计算 token 2 的注意力输出:
- 关注 token 1,2,3
- 假设 Q2 与 K1,K2,K3 的点积分别为 1, 2, 1
- 缩放后 softmax 得到权重,再对 V1,V2,V3 加权求和。
这样每个 token 只处理 3 个 key,而不是 6 个,计算量减少。
现代滑动窗口注意力的混合范式
概述
早期滑动窗口只是全局注意力的廉价替代。现代做法是全局注意力与局部/线性注意力交替:
- 例如每 4 层 Transformer 块为一组:
- 1 层使用完全全局注意力(看到所有历史 token)
- 3 层使用滑动窗口注意力(只看局部)
代表模型:Cohere Command A、LLaMA 4、Gemma 4、OLMo 3、Qwen 3.5 等。
数值例子(简化两层)
设序列长度 4,两层 Transformer:
- 第 1 层:全局注意力
- 第 2 层:滑动窗口注意力(窗口大小 2)
第 1 层全局注意力:
每个 token 都能看到所有 token,信息充分混合。
例如 token 3 可以看到 token 0,1,2,3。
第 2 层滑动窗口注意力:
每个 token 只看自己和前一个 token。
例如 token 3 只看 token 2,3。
信息传播:
虽然第 2 层只看局部,但第 1 层已经将全局信息融合到每个 token 的表示中。因此,经过多层交替,局部信息逐层向上汇聚,最终也能捕捉长距离依赖,同时计算成本大幅降低。
总结
| 机制 | 核心思想 | KV Cache 头数 | 优点 | 典型模型 |
|---|---|---|---|---|
| MHA | 每个头独立 Q,K,V | h | 表达能力强 | 原始 Transformer |
| MQA | 所有头共享 K,V | 1 | KV Cache 最小,推理最快 | PaLM, Falcon |
| GQA | 分组共享 K,V | G | 平衡效率与性能 | LLaMA 2/3, Qwen |
| 滑动窗口 | 只关注局部窗口 | - | 计算量线性增长 | Longformer, GPT-3 |
| 混合范式 | 全局与局部交替 | - | 长上下文高效 | Cohere Command A, LLaMA 4 |
一句话: MQA 共享全部 K,V,GQA 分组共享,二者通过减少 KV 头数提高算术强度、加速推理;滑动窗口注意力限制关注范围,混合范式交替使用全局与局部注意力,在长上下文场景中取得效率与性能的平衡。