Skip to content

高效注意力

你在 从RNN到Transformer 里已经见过:注意力的 n×n 矩阵意味着 O(n²) 的计算与显存。序列一长,成本指数级飙升。这一页讲现代大模型如何「驯服」这个二次方。

1. 问题到底出在哪

序列长度 n → 注意力矩阵大小 n×n → 显存/算力 O(n²)
   n=1k   : 1M 个权重/头       还行
   n=32k  : ~10亿 个/头         显存爆炸
   n=128k+: 更大                必须优化才能跑

关键洞察:真正占大头的有两块——

  • 注意力矩阵本身(n²,但可以不显式存)。
  • KV Cache(自回归生成时缓存的历史 Key/Value,随层数、头数、长度增长)。

2. FlashAttention:不显存 n² 矩阵(工程革命)

FlashAttention 不改变注意力的数学结果,只改变计算方式

传统:  完整算出 n×n 的 scores 矩阵 → 存下来 → 再 softmax → 再乘 V   (显存 O(n²))
Flash: 把 Q/K/V 分块(tiling)搬进 GPU 高速 SRAM, 边算边用 online-softmax 累加,
       从不把完整 n×n 矩阵写回显存                                    (显存 O(n))
   ┌──────────┬───────────┬────────────┬───────────┐
   │          │  结果等价  │   更省显存   │  更快      │
   ├──────────┼───────────┼────────────┼───────────┤
   │ 朴素实现  │   基准     │   O(n²)     │  基准      │
   │ FlashAttn│ ✅完全一致 │  O(n) ✅    │ 2~4× ✅   │
   └──────────┴───────────┴────────────┴───────────┘

它是现代长上下文模型的默认底座。PyTorch 里甚至一行就有:

python
import torch.nn.functional as F
# 内置融合注意力(is_causal=True 自动加因果掩码), 底层即 flash/内存高效实现
out = F.scaled_dot_product_attention(q, k, v, is_causal=True)

回想 搭建NanoGPT 里我们正是用了 scaled_dot_product_attention——你已经在享受 FlashAttention 的红利了。

3. MQA / GQA:给 KV Cache 瘦身(省推理显存)

标准多头里,每个头都有自己的一套 K、V,生成时全要缓存。MQA/GQA 让多个 Query 头共享较少的 K/V 头:

MHA(标准):  Q头=32  K头=32  V头=32    每头独立
MQA       :  Q头=32  K头=1   V头=1     所有Q共享一套KV → KV Cache 缩到 1/32, 但质量略降
GQA(折中)  :  Q头=32  K头=8   V头=8     分组共享 → 接近 MHA 质量, 显著省显存 ✅主流
KV Cache 大小 ∝ 层数 × KV头数 × 每头维度 × 序列长度
              └── GQA 正是砍这里的"KV头数" ──┘

LLaMA-2 70B、LLaMA-3、Qwen 等大量模型都用 GQA。

4. 滑动窗口 / 稀疏 / 线性注意力:从根上降复杂度

方法思路复杂度代表
局部/滑窗注意力每个词只注意附近 w 个O(n·w)Mistral 滑动窗口
稀疏注意力只算部分配对(如 BigBird)近似次线性Longformer、BigBird
线性注意力用核技巧改写,去掉 n²O(n)Performer、RETRO、部分新架构
状态空间模型干脆不用注意力(Mamba)O(n)Mamba、RWKV
全局注意力:  每词看全部 → 强但贵 O(n²)
局部(滑窗):  每词看邻居 → 便宜 O(n·w), 但远距离要靠堆层间接传递

5. 长上下文的组合拳

现代 128k~1M 上下文,通常多项叠加

RoPE(见下一页) + 位置插值/YaRN   →  把位置外推到训练没见过的长度
FlashAttention                    →  让 n² 计算与显存可控
GQA                               →  让 KV Cache 撑得住长序列
滑动窗口/稀疏                      →  进一步砍远程开销

小结

  • 注意力的软肋是 O(n²)(矩阵)与 KV Cache(生成)。
  • FlashAttention 用分块+online-softmax,结果不变但显存 O(n)、速度更快。
  • GQA/MQA 共享 K/V 头,给 KV Cache 瘦身,是长文本/大 batch 的关键。
  • 更激进的方向是滑窗/稀疏/线性注意力乃至非注意力架构(Mamba)。
  • 长上下文靠 RoPE + 插值 + Flash + GQA 的组合拳实现。

下一步

高效注意力解决了「算得快」,长上下文的另一半秘密在位置编码 → 位置编码演进RoPE