6. 理解 GQA
GQAGrouped Query Attention分组查询注意力是 Google 在 2023 年论文《GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints》中提出的一种 Attention 优化方法。它的核心目标是在尽量不损失模型效果的前提下大幅降低 KV Cache 的内存占用和推理成本。 (ACL Anthology)如果你已经理解了 Transformer 的 Multi-Head AttentionMHA那么 GQA 本质上就是让多个 Query Head 共享同一组 Key 和 Value Head。先理解问题为什么要优化 Attention标准 Transformer 的 AttentionA t t e n t i o n ( Q , K , V ) s o f t m a x ( Q K T d ) V Attention(Q,K,V) softmax\left(\frac{QK^T}{\sqrt d}\right)VAttention(Q,K,V)softmax(dQKT)V对于每个 Head都有独立的QueryKeyValue假设hidden_size 4096 num_heads 32那么Q: 32个头 K: 32个头 V: 32个头结构如下Head1: Q1 K1 V1 Head2: Q2 K2 V2 ... Head32: Q32 K32 V32这就是MHA (Multi Head Attention)MHA 最大的问题在训练阶段问题不大。但在推理阶段尤其生成长文本模型要保存历史 Token 的K Cache V Cache例如上下文长度 128K Head数 32KV Cache 会非常巨大。实际上大模型推理时最大的瓶颈之一就是 KV Cache。 (IBM)第一个解决方案MQAGoogle 之前提出MQA (Multi Query Attention)思路极其简单保留多个 Query HeadQ1 Q2 ... Q32但是所有 Head 共用一个 K 和 VK_shared V_shared变成Q1 ─┐ Q2 ─┤ Q3 ─┤ ... ├── K_shared Q32─┘ V_shared即32个Q 1个K 1个VKV Cache 直接缩小32 → 1 32 \rightarrow 132→1理论上节省32 倍 32倍32倍的 KV 存储。推理速度暴涨。 (ACL Anthology)但 MQA 有副作用虽然快了Q1 Q2 Q3 ... Q32都在使用同一个K V很多 Head 的表达能力被压缩了。效果通常会下降MHA GQA MQA论文发现MQA 推理速度很好但模型质量会有明显损失。 (ACL Anthology)GQA 的核心思想Google 想MQA太极端 MHA太昂贵于是取中间方案Grouped Query Attention假设32个Query Head不要32个KVMHA也不要1个KVMQA而是8个KV例如32个Q 8个KV每4个 Query Head 共用一个 KV Head。结构变成Q1 Q2 Q3 Q4 ↓ KV1 Q5 Q6 Q7 Q8 ↓ KV2 ... Q29 Q30 Q31 Q32 ↓ KV8这就是Group Size 4数学表示设h q Q u e r y H e a d s h_q Query HeadshqQueryHeadsh k v K V H e a d s h_{kv} KV HeadshkvKVHeads那么g r o u p s i z e h q h k v group\ size \frac{h_q}{h_{kv}}groupsizehkvhq例如32 / 8 4 32/8 432/84即4个Q共享1个KV三种 Attention 对比假设32个Query HeadMHA32Q 32K 32V结构Q1 - K1,V1 Q2 - K2,V2 ... Q32-K32,V32KV Cache100%效果最好MQA32Q 1K 1V结构Q1 Q2 ... Q32 共享K,VKV Cache1/32效果下降明显GQA32Q 8K 8V结构4个Q共享1个KVKV Cache1/4效果接近MHA所以MHA ←→ GQA ←→ MQAGQA 本质上是MHA 与 MQA 之间的折中方案为什么 GQA 特别适合大模型因为推理时KV Cache 大小约为O ( L × h k v × d ) O(L \times h_{kv} \times d)O(L×hkv×d)其中L 上下文长度h k v h_{kv}hkv KV Head 数量所以如果32 KV Heads → 8 KV Heads则KV Cache减少 75 减少75%减少75对于32K 64K 128K长上下文模型收益巨大。这也是为什么Llama 2 70BLlama 3MistralQwen 2DeepSeek等现代大模型广泛采用 GQA。从代码角度理解传统 AttentionQ.shape[B,L,32,d]K.shape[B,L,32,d]V.shape[B,L,32,d]GQAQ.shape[B,L,32,d]K.shape[B,L,8,d]V.shape[B,L,8,d]然后Krepeat_interleave(K,4)Vrepeat_interleave(V,4)逻辑上扩展成32heads供 Attention 使用。例如KV1 → Q1,Q2,Q3,Q4 KV2 → Q5,Q6,Q7,Q8...从信息论角度理解你之前问过RoPE 那些频率设计是不是拍脑袋GQA 反而比 RoPE 更容易理解。GQA 背后的观察其实是Query Head 差异很大不同 Head 学习不同模式语法 实体 位置 代码 推理因此Q需要保留多样性K/V Head 差异没那么大论文实验发现很多 Head 的 Key/Value 存在冗余。因此Q保持32个 KV压缩到8个性能损失很小。 (Hugging Face)一句话总结GQAGrouped Query Attention可以理解成保留大量 Query Head 的表达能力让多个 Query Head 共享较少数量的 Key/Value Head从而大幅减少 KV Cache占用更少显存同时保持接近 Multi-Head Attention 的效果。(ACL Anthology)如果后面你准备从零实现 Llama/Qwen我还可以继续讲GQA 的完整数学推导Llama3 中num_heads32, num_kv_heads8的具体实现PyTorch 版 GQA 源码逐行解析KV Cache 为什么能从 GQA 中获得巨大收益FlashAttention 与 GQA 的关系。