从经典self-Attention 到 Flash Attention:为什么我们「不必算出每一个 âᵢ」

10 阅读6分钟

写在前面

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)):

standard-attention.png

一句话版(精简路径):

path-summary.png

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 ]

整条路径可以记成:

x → (q, k, v) → a(分数)→ â(权重)→ o(输出)

对 (o_A, o_B, o_C) 同理。序列长度为 (N) 时,注意力权重在概念上就是一张 (N \times N) 的表。


二、符号别混:(a_i) 和 (\hat{a}_i)

符号名称阶段是否已归一化
(a_i)attention scoreSoftmax 前
(\hat{a}_i)attention weightSoftmax 后是,和为 1

Flash Attention 相关讨论里常问的那句:

我们真的需要算出每一个 attention weight (\hat{a}_i) 吗?

指的就是 Softmax 之后的那一组 (\hat{a})。


三、按图实现,为什么会慢、会爆显存?

数学上 Attention 很干净,慢往往出在 实现怎么碰显存

3.1 GPU 上的两层「仓库」

名称比喻特点
HBM大仓库容量大,读写相对慢
SRAM工作台极快,但很小

算力很强时,若实现不断在 HBM 与 SRAM 之间搬一张巨大的中间表,时间会耗在 搬运 上,而不是乘法上。这类情况叫 memory-bound

3.2 朴素实现在干什么?

对照上面的图,常见写法近似是:

  1. 算完所有 (q_i^\top k_j),得到完整 score 矩阵(约 (N \times N))
  2. 写回 HBM
  3. 再读回来做 Softmax
  4. 再读一遍,去乘 (V)

序列稍长,这张表就非常大。你要的其实只是每行对应的一个 (o) 向量,中间却为整张 (\hat{a}) 矩阵付了多次读写账单。

3.3 和「公式对不对」无关

  • 图上的公式是对的
  • 慢的是 「先物化整张权重再乘 V」 这条实现路径
  • Flash Attention 不修改 Attention 的数学定义,只改 计算顺序与中间结果是否落地

四、Flash Attention:不必写出每一个 (\hat{a}_i)

flash-attention.png

一句话版:

分块计算 + 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」等价。

于是可以:

算本块 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)正确同样正确

五、对比表:到底省在哪里?

维度普通 AttentionFlash Attention
是否写出完整 (N\times N) 权重通常要尽量不要
主要瓶颈HBM 读写大幅减少读写
数值结果标准 AttentionExact(非近似)
能否轻松画出完整 attention map容易基本拿不到(被「算没了」)
长序列显存随 (N^2) 压力大友好得多
短序列有时差不多加速不一定明显

常见误解:

  1. Flash Attention ≠ 近似 Attention(稀疏、低秩那类)。它是 IO 友好的精确算法。
  2. Flash Attention ≠ KV Cache。前者管「这一次怎么算」;后者管「生成时历史 K/V 不要重算」。

六、和 KV Cache 的分工

技术主要解决什么
Flash Attention单次 Attention 算得快、少写中间大矩阵
KV CacheDecode 时历史 K/V 复用,避免逐步重算
MQA / GQA / MLA让 KV Cache 本身更小

Prefill 长 prompt 时 Flash 很有用;Decode 则强依赖 KV Cache。两者常一起用,但不是一回事。


七、工程上是不是「改一个参数就行」?

对使用方来说,经常是改实现开关,但不是无条件生效。

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
回退环境不支持时可能静默退回普通实现,需自己看耗时/显存
sdpaPyTorch 自带,依赖少,很多场景已够用

业务开发多数是在换 backend,不是改网络结构。


八、三条直觉收束

  1. 图是定义:(q) 对所有 (k) 打分 → Softmax → 加权 (v) 得到 (o)。
  2. 慢在实现:为中间 (N\times N) 权重矩阵付出了大量 HBM 读写。
  3. Flash 的答卷:分块 + Online Softmax,在 SRAM 里直接累加 (o),不必为 Softmax 把每一个 (\hat{a}_i) 完整写回大仓库;数学结果不变,IO 账单大减。

下次再看到「Flash Attention 加速 Attention」,可以翻译成:

不是换了一种更模糊的注意力,而是 换了一种更省搬运的精确算法


九、小结

标准 Attention 图告诉我们「要算什么」;Flash Attention 告诉我们「可以不算完整的 (\hat{a}) 表,也能得到同一个 (o)」。抓住 memory-bound → 少写大矩阵 → online softmax 保证等价,这张图和这项优化就算真正接上了。