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

写在前面

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 朴素实现在干什么?

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

  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)

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) 压力大 友好得多
短序列 有时差不多 加速不一定明显

常见误解:

  1. Flash Attention ≠ 近似 Attention (稀疏、低秩那类)。它是 IO 友好的精确算法。
  2. 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,不是改网络结构。


八、三条直觉收束

  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 保证等价,这张图和这项优化就算真正接上了。

相关推荐
不一样的少年_1 小时前
我让 AI Agent 先别改代码,它怎么还是动手了?
人工智能·agent·ai编程
m0_640602441 小时前
2026 年餐饮收银系统前后端技术实现——核心架构与原理详解
后端·微服务·云原生·架构
Scene2161 小时前
Flux 与 Mono:Project Reactor 核心响应式类型深度解析
后端
lingran__1 小时前
C++ STL unordered系列(哈希) 底层剖析与模拟实现万字详解 | 基于哈希表,复刻 SGI-STL 泛型哈希容器架构
开发语言·c++·后端·哈希算法·哈希表·泛型编程·unordered系列
小林ixn1 小时前
NestJS 入门实战:从 0 到 1 撸一个 Todo CRUD,感受装饰器与模块化的优雅
后端·mvc·nestjs
小虎AI生活1 小时前
从"调教智能体"到"直接召唤专家":WorkBuddy 的专家/技能/专家团完整玩法
ai编程
阿弱1 小时前
graph-core 的边与命令模式设计
java·后端·agent
智驭未来掌门人1 小时前
利用Qt设计实现一款桌面程序
后端
leeyi1 小时前
Callback 源码:aspect_inject 切面注入(第87篇-E73)
aigc·agent·ai编程