Transformer 架构

注意力变体:FlashAttention / MQA / GQA 为什么出现

标准自注意力复杂度 O(n²),显存和速度都吃不消长序列。为降本提速,派生出 FlashAttention(IO 感知)、MQA/GQA(共享 KV)、稀疏/线性注意力。本文讲清每种解决什么、代价是什么。

一句话回答

标准注意力的瓶颈是 O(n²) 的显存与算力。三条优化路线分别打不同痛点:FlashAttention 解决「显存带宽」、MQA/GQA 解决「KV Cache 太大」、稀疏/线性注意力 解决「n² 计算本身」。

标准注意力的瓶颈

对长度 n、维度 d:每个头的注意力要算 n×n 的分数矩阵,复杂度和显存都随 n² 增长。

显存: O(n²·h)    算力: O(n²·d)

更致命的是显存带宽:标准实现要把巨大的 n×n 中间矩阵写回 HBM 再读回来算 softmax,GPU 大部分时间在「搬数据」而非「算」。长上下文(32k/128k)下这事直接炸。

标准:先算完整 N×N 分数矩阵(巨大、写回显存) N×N 矩阵 → 写回 HBM → 读回算 softmax 带宽瓶颈 GPU 干等

FlashAttention:不写出完整矩阵

FlashAttention(Dao 2022)的核心是 IO-aware:把 Q/K/V 切块(tiling),在 GPU 高速缓存(SRAM)里一块一块算 softmax 并累加,全程不把 N×N 中间矩阵写回显存。需要重算(recomputation)反向时的中间值,用「算力换带宽」。

收益:显存从 O(n²) 降到 O(n),速度数倍提升,且数值完全等价于标准注意力(不是近似)。是长上下文和高效训练推理的基石,v2/v3 进一步优化。
FlashAttention 不改数学结果,只是更聪明的「搬运+计算」顺序。它和下面的 MQA/GQA 是正交的——可以叠加。

MQA(Multi-Query Attention):共享 KV

标准 MHA 里,每个注意力头有自己的一套 Q、K、V。推理时每个 token 都要缓存一份 K、V(KV Cache),头数一多,缓存巨大、且读取 KV 的显存带宽成为瓶颈。

MQA(Shazeer 2019):所有头共享同一份 K 和 V,只保留多套 Q。KV Cache 直接缩小 h 倍。

MHA:每头各一份 KV → 缓存大 Q₁K₁V₁ Q₂K₂V₂ Q₃K₃V₃ … MQA:所有头共用一份 KV Q₁ Q₂ Q₃ … │ 共享 K V MHA 缓存 = h 份 MQA 缓存 = 1 份
代价:质量略降、训练可能不稳;但推理提速明显,PaLM 等采用。

GQA(Grouped-Query Attention):折中方案

GQA(Ainslie 2023)是 MHA 和 MQA 的折中:把 h 个头分成 g 组,每组共享一份 K、V。g=1 就是 MQA,g=h 就是 MHA。

Llama 2/3、Mistral、Qwen 等几乎都用 GQA。它几乎不掉点,却把 KV Cache 和推理延迟压下来一大截,是「质量-效率」的最优平衡点。一般取 g = 头数 / 若干(如 8 头分 4 组)。

稀疏 / 线性注意力:砍掉 n² 计算

前面几种大多只优化常数。若要真正摆脱 O(n²),得限制「每个 token 看多少个」:

稀疏注意力
只算局部窗口 + 少量全局 token(Longformer、BigBird),复杂度降到 O(n·√n) 或 O(n·w)。
线性注意力
用核函数把 softmax 重排,先算 V·Kᵀ 再乘 Q,复杂度 O(n·d²),与 n 线性(Linear Attention、RetNet、Mamba 类)。
代价:稀疏/线性会损失全局交互能力,质量或长程依赖上通常弱于稠密注意力,工业落地不如 FlashAttention+GQA 普遍。

对比一览

方案打的点KV 头数显存/速度质量是否近似
MHA(标准)—h基准最佳精确
FlashAttention显存带宽h显存↓O(n) 速度↑精确精确
MQAKV 缓存1缓存↓h倍 快略降精确
GQAKV 缓存g缓存↓ 快近 MHA精确
稀疏/线性n² 计算h复杂度↓可能降近似

总结