写在前面
Transformer 里最吃显存、也最容易成为推理/训练瓶颈的,往往是 Attention。网上讲 Flash Attention 的文章很多,公式一上来就 tiling、online softmax,容易让人觉得「又是一篇只有结论没有直觉的优化文」。
这篇文章换一个角度:先把一张标准的 Attention 数据流图读透,再问一个很朴素的问题------
我们真的需要先算出每一个 attention weight (\hat{a}_i),再去乘 (v) 吗?
Flash Attention 的回答是:不需要。 你要的是最终的 (o),不是那张巨大的权重表。下面就顺着数据流图,把「定义 → 瓶颈 → 改法 → 落地」串起来。
一、一张图看懂标准 Attention
下面这张图对应常见的多 token Attention 数据流(以四个 token (A,B,C,D) 为例,计算 (o_D)):
text
o_D
↑
+
┌──────┬──────┬──────┬──────┐
│ × │ × │ × │ × │
│ â_A │ â_B │ â_C │ â_D │
└──┬───┴──┬───┴──┬───┴──┬───┘
↑ ↑ ↑ ↑
┌────┴──────┴──────┴──────┴────┐
│ Softmax │
└────┬──────┬──────┬──────┬────┘
↑ ↑ ↑ ↑
a_A a_B a_C a_D
↑ ↑ ↑ ↑
·········· 与 q_D 打分 ··········
↑ ↑ ↑ ↑
v_A k_A q_A v_B k_B q_B v_C k_C q_C v_D k_D q_D
↑ ↑ ↑ ↑ ↑ ↑ ↑ ↑ ↑ ↑ ↑ ↑
└──┴──┘ └──┴──┘ └──┴──┘ └──┴──┘
x_A x_B x_C x_D
1.1 每个 token 变成三路
| 符号 | 作用 |
|---|---|
| (v) | Value,真正被加权求和的内容 |
| (k) | Key,用来和 Query 比相似度 |
| (q) | Query,当前 token「在问」的向量 |
1.2 以 (o_D) 为例的三步
① 打分(Attention Score)
a_i = \\frac{q_D\^\\top k_i}{\\sqrt{d}} \\quad (i \\in {A,B,C,D})
图中的 (a_A, a_B, a_C, a_D) 是 Softmax 之前的分数(logits),尚未归一化。
② Softmax → Attention Weight
\\hat{a}_i = \\mathrm{softmax}(a)_i = \\frac{e\^{a_i}}{\\sum_j e\^{a_j}}
(\hat{a}_i \ge 0) 且 (\sum_i \hat{a}_i = 1)。这才是乘在 (v) 上的权重。
③ 加权求和
o_D = \\hat{a}_A v_A + \\hat{a}_B v_B + \\hat{a}_C v_C + \\hat{a}_D v_D
整条路径可以记成:
text
x → (q, k, v) → a(分数)→ â(权重)→ o(输出)
对 (o_A, o_B, o_C) 同理。序列长度为 (N) 时,注意力权重在概念上就是一张 (N \times N) 的表。
二、符号别混:(a_i) 和 (\hat{a}_i)
| 符号 | 名称 | 阶段 | 是否已归一化 |
|---|---|---|---|
| (a_i) | attention score | Softmax 前 | 否 |
| (\hat{a}_i) | attention weight | Softmax 后 | 是,和为 1 |
Flash Attention 相关讨论里常问的那句:
我们真的需要算出每一个 attention weight (\hat{a}_i) 吗?
指的就是 Softmax 之后的那一组 (\hat{a})。
三、按图实现,为什么会慢、会爆显存?
数学上 Attention 很干净,慢往往出在 实现怎么碰显存。
3.1 GPU 上的两层「仓库」
| 名称 | 比喻 | 特点 |
|---|---|---|
| HBM | 大仓库 | 容量大,读写相对慢 |
| SRAM | 工作台 | 极快,但很小 |
算力很强时,若实现不断在 HBM 与 SRAM 之间搬一张巨大的中间表,时间会耗在 搬运 上,而不是乘法上。这类情况叫 memory-bound。
3.2 朴素实现在干什么?
对照上面的图,常见写法近似是:
- 算完所有 (q_i^\top k_j),得到完整 score 矩阵(约 (N \times N))
- 写回 HBM
- 再读回来做 Softmax
- 再读一遍,去乘 (V)
序列稍长,这张表就非常大。你要的其实只是每行对应的一个 (o) 向量,中间却为整张 (\hat{a}) 矩阵付了多次读写账单。
3.3 和「公式对不对」无关
- 图上的公式是对的
- 慢的是 「先物化整张权重再乘 V」 这条实现路径
- Flash Attention 不修改 Attention 的数学定义,只改 计算顺序与中间结果是否落地
四、Flash Attention:不必写出每一个 (\hat{a}_i)
text
q_D
│
├──────────── 块1:K/V 的 {A,B} ────────────┐
│ 算局部 score → Online Softmax 修正 │
│ 立刻累加进 o(SRAM 内) │
│ ▼
│ o(部分)
│ │
├──────────── 块2:K/V 的 {C,D} ────────────┤
│ 更新 max / sum,修正旧累积 │
│ 继续累加进 o │
│ ▼
└─────────────────────────────────────► 最终 o_D
✗ 不会完整写出:â_A, â_B, â_C, â_D 整张表到 HBM
✓ 数学上与「先 Softmax 再 Σ â·v」等价
一句话版:
分块计算 + Online Softmax,在 SRAM 里直接累加出 (o),避免把 (N\times N) 的 attention 矩阵完整写入 HBM;结果与标准 Attention 等价(exact,不是近似)。
4.1 分块(Tiling)
工作台(SRAM)一次放不下全部 (K,V),就切成小块。例如算 (o_D) 时:先处理 ({A,B}),再处理 ({C,D}),在 SRAM 内合并贡献。
4.2 Online Softmax
Softmax 需要整行的全局 max 与指数和。分块时维护:
- 目前见过的 最大值 (m)
- 目前见过的 指数和 (s)
新块若带来更大的 score,就用新旧 max 的差 修正 已累加的 (o),并更新 (s),保证与「一次看完全行再 softmax」等价。
于是可以:
text
算本块 score → 立刻和本块 v 结合 → 累加进 o
整行 ({\hat{a}_i}) 不必完整写回显存。
4.3 和原图的对应
| 原图步骤 | 朴素实现 | Flash Attention |
|---|---|---|
| 算 (a_i) | 常整表落地 | 分块在 SRAM 内算 |
| Softmax 得 (\hat{a}_i) | 显式得到整行权重 | Online 完成,不完整落盘 |
| (\sum \hat{a}_i v_i) | 再读权重乘 V | 边算边累加进 (o) |
| 最终 (o_D) | 正确 | 同样正确 |
五、对比表:到底省在哪里?
| 维度 | 普通 Attention | Flash Attention |
|---|---|---|
| 是否写出完整 (N\times N) 权重 | 通常要 | 尽量不要 |
| 主要瓶颈 | HBM 读写 | 大幅减少读写 |
| 数值结果 | 标准 Attention | Exact(非近似) |
| 能否轻松画出完整 attention map | 容易 | 基本拿不到(被「算没了」) |
| 长序列显存 | 随 (N^2) 压力大 | 友好得多 |
| 短序列 | 有时差不多 | 加速不一定明显 |
常见误解:
- Flash Attention ≠ 近似 Attention (稀疏、低秩那类)。它是 IO 友好的精确算法。
- Flash Attention ≠ KV Cache。前者管「这一次怎么算」;后者管「生成时历史 K/V 不要重算」。
六、和 KV Cache 的分工
| 技术 | 主要解决什么 |
|---|---|
| Flash Attention | 单次 Attention 算得快、少写中间大矩阵 |
| KV Cache | Decode 时历史 K/V 复用,避免逐步重算 |
| MQA / GQA / MLA | 让 KV Cache 本身更小 |
Prefill 长 prompt 时 Flash 很有用;Decode 则强依赖 KV Cache。两者常一起用,但不是一回事。
七、工程上是不是「改一个参数就行」?
对使用方来说,经常是改实现开关,但不是无条件生效。
python
model = AutoModelForCausalLM.from_pretrained(
"你的模型",
attn_implementation="flash_attention_2", # 或 "sdpa"
torch_dtype=torch.bfloat16,
device_map="auto",
)
| 注意点 | 说明 |
|---|---|
| 依赖 | flash_attention_2 通常要装 flash-attn |
| 硬件 | 较新的 NVIDIA GPU(常见 Ampere 及以后)更稳 |
| 精度 | 多为 fp16 / bf16 |
| 回退 | 环境不支持时可能静默退回普通实现,需自己看耗时/显存 |
sdpa |
PyTorch 自带,依赖少,很多场景已够用 |
业务开发多数是在换 backend,不是改网络结构。
八、三条直觉收束
- 图是定义:(q) 对所有 (k) 打分 → Softmax → 加权 (v) 得到 (o)。
- 慢在实现:为中间 (N\times N) 权重矩阵付出了大量 HBM 读写。
- Flash 的答卷:分块 + Online Softmax,在 SRAM 里直接累加 (o),不必为 Softmax 把每一个 (\hat{a}_i) 完整写回大仓库;数学结果不变,IO 账单大减。
下次再看到「Flash Attention 加速 Attention」,可以翻译成:
不是换了一种更模糊的注意力,而是 换了一种更省搬运的精确算法。
九、小结
标准 Attention 图告诉我们「要算什么」;Flash Attention 告诉我们「可以不算完整的 (\hat{a}) 表,也能得到同一个 (o)」。抓住 memory-bound → 少写大矩阵 → online softmax 保证等价,这张图和这项优化就算真正接上了。