标准自注意力复杂度 O(n²),显存和速度都吃不消长序列。为降本提速,派生出 FlashAttention(IO 感知)、MQA/GQA(共享 KV)、稀疏/线性注意力。本文讲清每种解决什么、代价是什么。
标准注意力的瓶颈是 O(n²) 的显存与算力。三条优化路线分别打不同痛点:FlashAttention 解决「显存带宽」、MQA/GQA 解决「KV Cache 太大」、稀疏/线性注意力 解决「n² 计算本身」。
对长度 n、维度 d:每个头的注意力要算 n×n 的分数矩阵,复杂度和显存都随 n² 增长。
更致命的是显存带宽:标准实现要把巨大的 n×n 中间矩阵写回 HBM 再读回来算 softmax,GPU 大部分时间在「搬数据」而非「算」。长上下文(32k/128k)下这事直接炸。
FlashAttention(Dao 2022)的核心是 IO-aware:把 Q/K/V 切块(tiling),在 GPU 高速缓存(SRAM)里一块一块算 softmax 并累加,全程不把 N×N 中间矩阵写回显存。需要重算(recomputation)反向时的中间值,用「算力换带宽」。
标准 MHA 里,每个注意力头有自己的一套 Q、K、V。推理时每个 token 都要缓存一份 K、V(KV Cache),头数一多,缓存巨大、且读取 KV 的显存带宽成为瓶颈。
MQA(Shazeer 2019):所有头共享同一份 K 和 V,只保留多套 Q。KV Cache 直接缩小 h 倍。
GQA(Ainslie 2023)是 MHA 和 MQA 的折中:把 h 个头分成 g 组,每组共享一份 K、V。g=1 就是 MQA,g=h 就是 MHA。
前面几种大多只优化常数。若要真正摆脱 O(n²),得限制「每个 token 看多少个」:
| 方案 | 打的点 | KV 头数 | 显存/速度 | 质量 | 是否近似 |
|---|---|---|---|---|---|
| MHA(标准) | — | h | 基准 | 最佳 | 精确 |
| FlashAttention | 显存带宽 | h | 显存↓O(n) 速度↑ | 精确 | 精确 |
| MQA | KV 缓存 | 1 | 缓存↓h倍 快 | 略降 | 精确 |
| GQA | KV 缓存 | g | 缓存↓ 快 | 近 MHA | 精确 |
| 稀疏/线性 | n² 计算 | h | 复杂度↓ | 可能降 | 近似 |