从 MHA 到 GQA 到 MQA,Attention 结构的演进是推理优化的核心线索
前置知识
核心概念
为什么 Attention 头结构影响部署
KV Cache 是推理中最大的显存消耗。Attention 头数直接决定 KV Cache 大小。而 Query、Key、Value 的分组方式(MHA / GQA / MQA)决定了每个 token 需要缓存多少数据。
MHA 详细工作原理
输入: hidden_states [batch, seq_len, d_model] 其中 d_model = num_heads × head_dim
1. 线性投影: Q = hidden @ W_q → [batch, seq_len, num_heads × head_dim] K = hidden @ W_k → [batch, seq_len, num_heads × head_dim] V = hidden @ W_v → [batch, seq_len, num_heads × head_dim]
2. reshape 为多头: Q → [batch, seq_len, num_heads, head_dim] → transpose → [batch, num_heads, seq_len, head_dim] K → [batch, seq_len, num_heads, head_dim] → transpose → [batch, num_heads, seq_len, head_dim] V → [batch, seq_len, num_heads, head_dim] → transpose → [batch, num_heads, seq_len, head_dim]
3. Attention 计算: scores = Q @ K^T / sqrt(head_dim) → [batch, num_heads, seq_len, seq_len] attn = softmax(scores, dim=-1) output = attn @ V → [batch, num_heads, seq_len, head_dim]
4. 拼接 + 输出投影: output → transpose → [batch, seq_len, num_heads, head_dim] → reshape → [batch, seq_len, d_model] → @ W_o → [batch, seq_len, d_model]
以 Llama 3 8B 为例:
MHA / GQA / MQA 结构对比
GQA 的分组策略
GQA 的核心思想:将 num_heads 个 Query head 分成 G 组,每组共享一组 K 和 V。
num_kv_heads = num_q_heads / G
KV Cache 减少倍数 = G 倍
为什么 8 KV groups 是主流选择?
8 groups 是经验和实验得出的甜蜜点:
MQA 质量下降原因分析
MQA 将所有 Q heads 共享一组 KV,极端压缩了 KV Cache。但质量下降的原因在于:
实验数据:MQA 相比 MHA 通常有 1-3% 的质量下降,而 GQA-8 几乎无下降。
FlashAttention 原理简介
理解 FlashAttention 需要先了解 GPU 的内存层级:HBM(High Bandwidth Memory,GPU 的大容量显存,~80GB)和 SRAM(片上缓存,~192KB)。HBM 容量大但访问慢,SRAM 极快但容量极小。FlashAttention 的核心就是把计算搬到 SRAM 里做。
FlashAttention 是一种 IO-aware 的 Attention 实现,核心思路:
传统 Attention: QK^T 结果存到 HBM → softmax → 结果存到 HBM → 读取后与 V 相乘 HBM 访问次数: O(N^2) 次读 + O(N^2) 次写(N = seq_len)
FlashAttention: 将矩阵分块,每次只加载一个 tile 到 SRAM(片上缓存) 在 SRAM 内完成 QK^T → softmax → V 的完整计算 HBM 访问次数: O(N^2) 但常数大幅减小
关键优化: - 利用 SRAM(~192KB on A100)替代 HBM 的反复读写 - 在线 softmax(online softmax)避免存储完整的 attention matrix - 重计算(recomputation)反向传播时重算而非存储
效果:在 A100 上,FlashAttention 2 比标准 Attention 快 2-4x ,显存减少 ~20% 。
部署视角
GQA 对部署的具体影响
batch size 上限提升 :
假设 A100 80GB,预留 10GB 给权重和其他: 可用显存给 KV Cache: ~70GB
MHA (32 KV heads, Llama-1 13B 级别): 每 token 每层 KV: 2 × 32 × 128 × 2 = 16 KB 40 层, seq=4096: 16 KB × 40 × 4096 × batch / 1GB batch_max ≈ 16
GQA (8 KV heads): 每 token 每层 KV: 2 × 8 × 128 × 2 = 4 KB 同样条件下: batch_max ≈ 64
结论:GQA 将 batch size 上限提升 4-8 倍。
实际数字对比 :
KV Cache 节省的具体数字
Llama 3 70B, FP16, batch=32, seq_len=8192:
MHA (假设 32 KV heads): KV = 2 × 80 × 32 × 8192 × 32 × 128 × 2 ≈ 343 GB
GQA (8 KV heads): KV = 2 × 80 × 32 × 8192 × 8 × 128 × 2 ≈ 86 GB
节省: 343 - 86 = 257 GB (减少 75%)
这决定了 70B 模型能否在单张 A100 80G 上部署。
常见问题排查
面试视角
面试官会怎么问
Q1: "GQA 是怎么减少 KV Cache 的?为什么 8 groups 是主流?"
满分回答:
Q2: "FlashAttention 为什么比普通 Attention 快?"
Q3: "MQA 为什么能最小化 KV Cache?它有什么代价?"
Q4: "MHA 中 Q/K/V 的维度是怎么变换的?"
对比分析
三种 Attention 方案全面对比
最佳实践
调参建议
避坑指南