Appearance
高效注意力
你在 从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。