第三章:MHA 的局限与 MQA、GQA、Flash Attention¶
3.1 MHA 的瓶颈卡在哪¶
讲清楚 MHA 的局限,必须把训练和推理两个阶段分开看。很多人只笼统说「\(O(N^2)\) 慢」,一追问就答不下去,根源就是把这两个阶段混在一起讲。
3.1.1 训练阶段¶
长度为 \(N\) 的序列,每层 Attention 都要算一个 \(N \times N\) 的注意力分数矩阵,而且这个矩阵要存在显存里给反向传播用:
| 序列长度 | \(N \times N\) 矩阵元素数 | FP16 显存(单层单头) |
|---|---|---|
| 4K | 1600 万 | 约 32 MB |
| 32K | 10 亿 | 约 2 GB |
计算复杂度是 \(O(N^2)\),\(N\) 翻倍计算量翻 4 倍。
好消息是训练时这个矩阵是「一次性算完」的,不需要在多个时间步之间持续保留。
3.1.2 推理阶段:KV Cache 是救星也是显存大户¶
LLM 是自回归生成,每生成一个新 token 都要对前面所有 token 算注意力。如果每次从头算,总成本会累加到 \(O(N^3)\)。
KV Cache 的做法是把前面所有 token 的 K 和 V 存下来,新 token 只算自己的 Q,和缓存的 K/V 做注意力。这是推理优化的标配(详见 第十四章)。
但它本身很吃显存:
(前面的 2 是 K 和 V 各一份,后面的 2 是 FP16 每个数 2 字节)
以一个 7B 模型(\(L=32\)、\(H=32\)、\(d_k=128\))跑 \(B=1\)、\(N=32000\) 为例:
光 KV Cache 就 17GB,加上模型权重 14GB 一共 31GB——一张 24GB 的 4090 根本放不下。
3.1.3 更隐蔽的痛点:访存受限¶
显存挤爆只是一面,另一面是速度也快不起来。
GPU 的计算单元算力很猛,但显存带宽跟不上。Attention 计算里大量时间花在「等数据从显存搬到计算单元」,计算单元很多时候在「等米下锅」——这就是 memory-bound(访存受限)。
哪怕你的 GPU 算力是另一台的两倍,跑 Attention 时可能只快 20%,因为瓶颈根本不在算力。
3.1.4 三个痛点连成一条线¶
flowchart TB
P1["痛点一 · O(N²) 复杂度<br/>序列稍长计算量按平方膨胀"]
P2["痛点二 · KV Cache 显存<br/>长上下文直接吃光显存"]
P3["痛点三 · 访存带宽<br/>GPU 算力发挥不出来"]
P1 --> LONG["长上下文场景<br/>三者互相加剧"]
P2 --> LONG
P3 --> LONG
style LONG fill:#fce8e6
后面所有优化方案,都是在攻击这三个痛点中的一个或多个。
3.2 MQA:暴力共享 K/V¶
3.2.1 思路¶
所有 head 共享同一份 K 和 V,只有 Q 每个 head 独立。
原本 MHA 有 32 个 head,每个都有自己的 \(W_Q\)、\(W_K\)、\(W_V\),输出 32 套 Q、K、V。MQA 只保留 32 套 Q,K 和 V 全部 32 个 head 共享一套。
flowchart LR
subgraph MHA_D["MHA · H=4"]
Q1[Q1]---K1[K1/V1]
Q2[Q2]---K2[K2/V2]
Q3[Q3]---K3[K3/V3]
Q4[Q4]---K4[K4/V4]
end
subgraph MQA_D["MQA · H=4"]
MQ1[Q1] --> SKV[共享 K/V]
MQ2[Q2] --> SKV
MQ3[Q3] --> SKV
MQ4[Q4] --> SKV
end
style SKV fill:#e6f4ea
3.2.2 收益与代价¶
收益:KV Cache 直接变成 \(1/H\)。上面那个 17GB 的例子,用 MQA 只剩 0.5GB 多一点。
而且注意力公式不变——每个 head 还是各算各的分数,只是共用同一份 K,模型结构基本保持,训练流程几乎不用改。
代价是表达能力下降。
直觉理解:原本 32 个 head 各有 32 套不同的「视角」(K 表示「我有什么标签」、V 表示「我的内容」),可以从 32 个角度理解上下文。MQA 让它们共用一套,等于32 个视角变成「都看同样的标签和内容,只是用不同的 Query 去问」,多视角能力被压成单视角。
实测在大模型上效果下降 2–5%。简单任务差不多,但对推理类任务(数学、代码)损失明显。所以 MQA 在工业界不如它的折中版本受欢迎。
3.3 GQA:效果与显存的甜蜜点¶
3.3.1 思路¶
把 \(H\) 个 head 分成 \(G\) 组,每组内部共享一份 K/V,组之间各自独立。
| 方案 | head 数 | K/V 套数 |
|---|---|---|
| MHA | \(H\) | \(H\) |
| GQA | \(H\) | \(G\)(\(1 \le G \le H\)) |
| MQA | \(H\) | 1 |
GQA 是一个连续光谱:\(G=H\) 退化成 MHA,\(G=1\) 退化成 MQA,中间任意取值都行。
3.3.2 为什么是好折中¶
| 维度 | 效果 |
|---|---|
| 显存 | KV Cache 从 \(H\) 套压到 \(G\) 套,占用比例 \(G/H\)。\(H=32, G=8\) 时压到 1/4 |
| 表达力 | 每组有自己的「视角」,组数越多视角越丰富,\(G=8\) 通常已足够 |
| 实测 | GQA 论文中 \(G=8\) 配置下效果几乎和 MHA 持平(差距不到 0.5%) |
「显存大幅下降、效果几乎不损失」这个甜蜜点,让 GQA 成为现代大模型的标配:
| 模型 | 配置 |
|---|---|
| Llama 2 70B | GQA,\(H=64\)、\(G=8\) |
| Llama 3 全系 | GQA |
| Qwen 2 / 3 主力模型 | GQA |
| DeepSeek V2 / V3 | MLA(另一条路线,见下) |
3.3.3 MLA 不是 GQA 的升级版¶
DeepSeek V2/V3 用的 MLA(Multi-head Latent Attention,多头潜在注意力) 走的是另一条路:
不共享 K/V,而是把每个 token 的 K/V 通过降维投影压缩到低维 latent 向量存起来,需要时再配合额外投影参与注意力计算。
存的不是「\(H\) 套或 \(G\) 套高维 K/V」,而是「低维压缩后的表示」,显存比传统 MHA / GQA 更省。
目标相似(都是省 KV Cache),但实现机制完全不同。别说成「GQA 的升级版」——GQA 是「少存几套」,MLA 是「换一种更紧凑的表示存」。
3.4 Flash Attention:换一条赛道¶
MQA/GQA 改的是 Attention 结构(有几套 K/V)。Flash Attention 完全是另一条赛道——不改数学公式,从底层实现优化。
3.4.1 根源:显存层级的巨大差距¶
GPU 的存储分两层:
| 层级 | 容量 | 带宽 |
|---|---|---|
| HBM(高带宽显存,平时说的「显存」) | A100 是 40/80 GB | 约 1.5 TB/s |
| SRAM(片上缓存) | A100 每个 SM 仅 192 KB | 约 19 TB/s(HBM 的 13 倍) |
标准 Attention 的实现是这样的:
S = Q @ K.T # 算出 N×N 分数矩阵,写回 HBM
P = softmax(S) # 从 HBM 读 S,算 softmax,再写回 HBM
O = P @ V # 从 HBM 读 P,算最终输出,再写回 HBM
整个过程在 HBM 上反复读写 \(N \times N\) 的大矩阵,访存时间远超实际计算时间。
这就是为什么 \(N=4K\) 的注意力比 \(N=2K\) 慢的倍数往往超过理论上的 4 倍——瓶颈在于搬运 \(N^2\) 大小的中间结果。
3.4.2 核心思路:分块 + 在线 softmax¶
既然 SRAM 带宽快但容量小,那就把 Q、K、V 切成小块(比如 128×128),每次只在 SRAM 里算一小块的注意力,算完直接和最终输出累加,不把 \(N \times N\) 的中间矩阵写回 HBM。
flowchart TB
subgraph STD["标准 Attention"]
S1["算 N×N 分数矩阵"] --> S2["写回 HBM"]
S2 --> S3["读回算 softmax"] --> S4["写回 HBM"]
S4 --> S5["读回乘 V"] --> S6["写回 HBM"]
end
subgraph FA["Flash Attention"]
F1["切成 128×128 小块"] --> F2["整块搬进 SRAM"]
F2 --> F3["在 SRAM 里算完这一块"]
F3 --> F4["在线 softmax 增量更新<br/>直接累加到输出 O"]
F4 -->|下一块| F2
F4 --> F5["只把最终 O 写回 HBM"]
end
style STD fill:#fce8e6
style FA fill:#e6f4ea
这里有个数学难题:softmax 要看「整行」才能算,不能局部独立计算。
Flash Attention 用在线 softmax(online softmax)解决:分块计算的同时维护「当前最大值 + 累积和」的状态,每来一块做增量更新,最终结果和一次性算 softmax 完全一样。
3.4.3 三个收益¶
| 收益 | 说明 |
|---|---|
| 显存 \(O(N^2) \to O(N)\) | 不需要把中间矩阵存 HBM,只存最终输出 |
| 速度快 2–4 倍 | HBM 读写次数从 \(O(N^2)\) 降到 \(O(N^2/M)\)(\(M\) 是块大小) |
| 数学等价 | 算的是同一个公式,不是稀疏或低秩近似 |
关于「等价」要精确一点:实际浮点实现里因为分块顺序和数值精度不同,最后几位可能有微小差异,但不会像近似注意力那样引入模型精度损失。
Flash Attention 现已迭代到 v3,针对 H100 等新一代 GPU 做了进一步优化,是 vLLM、SGLang、TGI 等主流推理框架的默认实现。
3.5 三类优化是叠加不是替代¶
| 方案 | 改的是什么 | 攻击的痛点 | 效果损失 |
|---|---|---|---|
| MQA | 结构(K/V 压成 1 份) | 显存 | 中等(2–5%) |
| GQA | 结构(K/V 压成 G 份) | 显存 | 几乎无(不到 0.5%) |
| Flash Attention | 实现(分块 + 在线 softmax) | 显存 + 访存 + 速度 | 基本无(数学等价) |
关键在于:
flowchart TB
STRUCT["结构层优化<br/>GQA · 决定要存几套 K/V"]
IMPL["实现层优化<br/>Flash Attention · 决定怎么算"]
STRUCT --> COMBO["主流大模型标配<br/>GQA 结构 + Flash Attention 实现"]
IMPL --> COMBO
COMBO --> R["7B 模型在 4090(24GB)上<br/>跑 32K 长上下文"]
style COMBO fill:#e6f4ea
Llama 3、Qwen 2、DeepSeek V3 都是这个组合:GQA 把 KV Cache 压到 1/4,Flash Attention 让计算快 2–4 倍。
面试官如果问「MQA、GQA、Flash Attention 只能选一个,选哪个」,正确回答是指出这是个伪命题——它们攻击的是不同维度的瓶颈,真实工程里一定组合用。能说出这句,说明你理解的是整套优化体系的层次结构,不是在背单点优化。
3.6 长上下文时代的其他方向¶
再补几条常见路线:
| 方向 | 思路 | 现状 |
|---|---|---|
| MLA | K/V 压缩到低维潜在空间 | DeepSeek V2/V3 在用,已工程验证 |
| Sliding Window Attention | 每个 token 只关注最近 \(N\) 个(如 4096) | Mistral 系列用过,代价是丢失远距离信息,通常和全局注意力混合 |
| Linear Attention | 用核函数近似 softmax,复杂度降到 \(O(N)\) | Performer、Linformer 等,效果离 MHA 差一截,未成主流 |
| Mamba / SSM | 完全抛弃 Attention,用状态空间方程 | 理论上可处理无限长上下文,效果仍有争议,研究热点 |
面试深度建议:把 MQA、GQA、Flash Attention 三个讲透就足够拿高分,MLA 作为加分项提一句,再深的不用展开。
3.7 常见错误¶
3.7.1 只说「\(O(N^2)\) 慢」¶
这只是三个痛点之一,而且是最表面的一个。要能分训练/推理两阶段讲,说出 KV Cache 显存和访存受限这两个更工程化的痛点。
3.7.2 说不出 KV Cache 的量级¶
能现场算出「7B 模型 32K 上下文约 17GB」这个数,比空谈「显存很大」有说服力得多。
3.7.3 忽略访存受限¶
这是最容易被漏掉的一点,也是理解 Flash Attention 的前提。瓶颈不在算力在带宽——不知道这一点就理解不了为什么「不改公式只改实现」能快 2–4 倍。
3.7.4 认为 Flash Attention 是近似算法¶
它是数学等价的精确实现,不是稀疏近似或低秩近似。这是它能被无条件默认启用的原因。
3.7.5 把 MQA 和 GQA 当成两种独立方案¶
GQA 是一个连续光谱,两端分别退化成 MHA 和 MQA。
3.7.6 把 MLA 说成 GQA 的升级版¶
GQA 是「少存几套」,MLA 是「换一种低秩压缩表示存」,机制不同。
3.7.7 把三类优化当成互斥选项¶
结构层和实现层攻击不同维度,主流模型都是 GQA + Flash Attention 同时用。
3.8 本章总结¶
- 必须分训练和推理两阶段讲 MHA 的瓶颈,两者痛点不同;
- 三个痛点:\(O(N^2)\) 复杂度、KV Cache 显存、访存受限,长上下文场景下互相加剧;
- KV Cache 的量级要能算:7B 模型 32K 上下文约 17GB,加权重放不进 24GB 显卡;
- MQA 把 K/V 压成 1 套,显存降到 \(1/H\),但表达力损失 2–5%,推理类任务尤其明显;
- GQA 分 G 组共享,是连续光谱,\(G=8\) 时显存 1/4、效果差距不到 0.5%,现代大模型标配;
- MLA 是另一条路线:低秩压缩存储,不是 GQA 的升级版;
- Flash Attention 走实现层赛道:分块 + 在线 softmax,显存 \(O(N^2) \to O(N)\),速度 2–4 倍,且数学等价;
- 理解 Flash Attention 的前提是知道 HBM 与 SRAM 的 13 倍带宽差距;
- 三类优化是叠加不是替代,GQA + Flash Attention 是主流组合,让 7B 模型能在消费级显卡上跑 32K。
参考资料¶
- Attention Is All You Need
- Fast Transformer Decoding: One Write-Head is All You Need(MQA)
- GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints
- FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness
- FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning
- DeepSeek-V2: A Strong, Economical, and Efficient Mixture-of-Experts Language Model(MLA)
- Online normalizer calculation for softmax
- Mistral 7B(Sliding Window Attention)
- Mamba: Linear-Time Sequence Modeling with Selective State Spaces