标准 Attention 的计算流程是 、、,其中 ( 为序列长度, 为头维度)。两次矩阵乘法的总 FLOPs 为 ,但硬件真实瓶颈往往不在计算量,而在中间矩阵 和 :它们各占 的显存,必须完整写入 HBM 再读出。这不仅导致 IO 成为性能瓶颈(算术强度1约 FLOPs/Byte,远低于 A100 拐点 156),更直接带来显存容量问题:、32 个头时,仅 和 就要占用约 2 GB 显存,序列再长就直接 OOM。
FlashAttention 的优化方向因此是减少 HBM 访存,同时消除 矩阵的显存占用。做法是两个技术的配合:
- Tiling 将 切成小块在 SRAM 中处理,避免 落地 HBM2
- Online Softmax 解决 softmax 这个非线性操作在分块计算时的归一化问题:标准 Softmax 需要收集最大值、指数和,而分块时一次只能看到局部数据。Online Softmax 用递推公式动态修正局部结果,使最终输出与标准 Attention 数学上完全等价32。
本文从这两个基础技术出发,逐步拆解 FlashAttention V1 和 V2 的计算过程。如果你对 Attention 基本原理还不熟悉,建议先阅读《为什么 KV 缓存没有 Q》。
FlashAttention 系列总览
FlashAttention 从 2022 年至今迭代了三个版本,核心思路一脉相承:将 Attention 计算拆成小块,在 SRAM 中完成,避免 中间矩阵落地 HBM。
V1245(Dao et al., 2022)引入 Tiling + Online Softmax:
- 将 沿序列维度分块,每次只在 SRAM 中计算一小块
- 用 Online Softmax 的递推公式实现精确的分块归一化
- 反向传播不存储 ,从少量统计量(行最大值 、指数和 )出发重计算
- 额外显存从 降至 ,HBM 访问量降至 ,证明为渐近最优2
- A100 上前向加速 2-4x,GPT-2 训练端到端加速约 3x2

V267(Dao, 2023)做了三项工程优化,将 A100 利用率从 25-40% 提升到 50-73%:
- 翻转循环顺序:Q 块做外循环常驻 SRAM,K/V 块流过。 只写回 HBM 一次(V1 是读写共 次),且 Q 块天然映射到不同 thread block
- 延迟归一化:维护未归一化的累加器 ,最后一步才除以 ,减少每步非 GEMM 操作
- Warp 沿 Q 维度切分:每个 warp 独占自己的行,无需跨 warp 同步
V38(Shah, Dao et al., 2024)针对 H100 Hopper 架构做硬件级适配:
- Warpgroup 间 pingpong 调度:两个 warpgroup 交替执行 GEMM 和 softmax,隐藏非 GEMM 延迟
- Warpgroup 内流水线重叠:利用 Hopper 异步执行能力,GEMM 和 softmax 的不同阶段并行
- FP8 + incoherent processing:block quantization 与 incoherent processing 共同降低 FP8 量化误差,其中后者用随机正交变换(实践中包含 Hadamard 变换)打散异常值,误差降低 2.6x8
下面先给出 Online Softmax 和 Tiling 的公式框架,再进入 V1 和 V2 的完整算法。
两个核心技术
Online Softmax:流式精确归一化
标准 Safe Softmax 对输入向量 (在 Attention 语境下即 的某一行)需要三遍扫描:找全局最大值 (避免后续指数溢出),算指数和 ,算归一化输出。三遍之间有严格的数据依赖。分块计算时无法提前获得全局 ,Online Softmax3 的解法是边扫描边用递推公式动态修正。
核心恒等式:
发现更大的 时,之前以 为基准的指数项乘以修正因子 即可对齐。纯代数恒等变形,没有近似。
设 被分成 块,上标 表示处理完前 个块后的累积状态。当第 块到来时:
从 Softmax 到 Attention Output
Attention 最终需要的不是 Softmax 概率本身,而是 。Online Softmax 的递推直接扩展到输出 上。设 和 分别为第 个 K/V 块对应的注意力分数和 Value,处理完前 个块后,未归一化的输出累加器为:
所有块处理完毕后一次性归一化:。

Tiling:分块驻留 SRAM
沿序列维度分成 块, 分成 块。块大小由 SRAM 容量 (元素数)和头维度 决定:
分母 4 可以理解为论文在 SRAM 预算下给出的保守容量约束:片上主要部分是 四类 tile。实际实现里 和 可以不同,也会受中间状态、统计量和硬件资源占用影响。

FlashAttention V1:算法详解
前向传播
V1 的前向传播是一个双层循环,外循环遍历 K/V 块(),内循环遍历 Q 块():
初始化 O = 0, m = -∞, l = 0
外循环:for j = 1 to T_c
从 HBM 加载 K_j, V_j 到 SRAM(常驻至内循环结束)
内循环:for i = 1 to T_r
从 HBM 加载 Q_i, O_i, m_i, l_i 到 SRAM
计算 S_ij = Q_i · K_j^T / √d ← SRAM 内完成
更新 m, l, O(Online Softmax) ← SRAM 内完成
将 O_i, m_i, l_i 写回 HBM
每次内循环迭代的计算步骤:
Step 1 — 局部注意力分数:
Step 2 — 更新行最大值:
Step 3 — 更新指数和:
Step 4 — 更新输出:
几个细节:
初始化为 , 初始化为 。 第一块数据进来时 ,修正项自动清零,公式退化为普通的局部计算。用数学边界值替代 if (is_first_iteration) 分支判断。
和 是向量()而非标量。 Softmax 按行独立归一化: 的每一行对应一个 Query token,各有自己的注意力分布。 个 token 就需要 个独立的 和 。
数值稳定性。 所有指数输入都减去了当前最大值,保证 。修正因子 ,不会溢出。
CUDA 执行模型
V1 以 grid = (Batch, Head) 启动内核:每个 (batch, head) 对应一个 thread block,被调度到某个 SM 上执行(一个 SM 可同时驻留多个 thread block,具体取决于资源占用),串行跑完双层循环。
- Grid 级别(thread block 间):完全并行,无通信
- Block 级别(循环迭代间):严格串行,不存在 和 同时计算的情况
- Thread/Warp 级别(单次迭代内):多线程协作完成矩阵乘法和归约操作
瓶颈:并行度绑定在 batch size head 数上。batch=1、head=32 时,A100 的 108 个 SM 只有 32 个在工作。这是 V2 翻转循环顺序的动因。
反向传播与重计算
FlashAttention 不存储 和 ,前向只保存 ,反向时在 SRAM 中重计算:
重计算增加约 50% 前向 FLOPs,但由于标准 Attention 是 memory-bound 的,额外计算时间远小于省下的 IO 时间。以 A100 处理 的单头为例:
| 项目 | 量 | A100 耗时 |
|---|---|---|
| 重计算 | FLOPs | |
| 从 HBM 读 矩阵 | MB |
片上重算比从 HBM 读一遍还快,而标准反向传播涉及 的多次读写。重计算用微量计算时间换下大量 IO 等待,墙钟时间反而缩短。
IO 复杂度
标准 Attention 的 HBM 访问量为 , 项来自中间矩阵 。
FlashAttention V1 的 HBM 访问量为 2。推导:外循环 次迭代,每次读取全部 (共 ),总计 。由于 ( 时 ,A100 SRAM 约 个 FP16 元素),这远小于 。
精确 Attention 的 HBM 访问下界为 ,因此 FlashAttention 的 IO 复杂度是渐近最优的2。
FlashAttention V2:三项改进
V1 在 A100 上只达到 25-40% 理论算力利用率。V2 将其提升到 50-73%,比 V1 快约 1.7x6。
翻转循环顺序
V1 外循环 K/V、内循环 Q,每次外循环切换 时所有 Q 块的统计量都要重新从 HBM 加载再写回, 被读写 次。
V2 反转:外循环 Q 块(映射到不同 thread block 并行),内循环 K/V 块。
外循环(并行):for i = 1 to T_r
Q_i 加载到 SRAM(全程常驻)
初始化 m_i = -∞, l_i = 0, Õ_i = 0
内循环:for j = 1 to T_c
从 HBM 加载 K_j, V_j 到 SRAM
计算 S_ij,更新 m, l, Õ(全在 SRAM 中)
O_i = Õ_i / l_i ← 最终归一化
L_i = m_i + log(l_i) ← 保存 logsumexp
将 O_i, L_i 写回 HBM ← 只写一次
三个收益: 只写 HBM 一次;Q 全程驻留 SRAM;不同 Q 块之间无递推依赖,Grid 从 (Batch, Head) 扩展为 (Batch, Head, T_r)。
延迟归一化
V1 每步维护归一化后的 ,需要先反归一化(乘回旧分母)再重归一化(除以新分母),每步 3 次标量运算。V2 改为维护未归一化的 :
最后一步才归一化:。这些逐元素运算跑在 CUDA Core 上,而矩阵乘法跑在 Tensor Core 上。A100 的 FP16/BF16 matmul 峰值约 312 TFLOPS,non-matmul FP32 只有 19.5 TFLOPS,单个 non-matmul FLOP 约贵 16 倍6。减少非 GEMM 操作直接降低了这部分的时间占比。
Warp 分工沿 Q 维度切分
V1 沿 K/V 维度切分 warp,每个 warp 只拿到 的部分列,softmax 的 rowmax/rowsum 必须跨 warp 同步合并。V2 沿 Q 维度切分,每个 warp 独占自己的行,无需通信:
| V1(沿 K/V 切分) | V2(沿 Q 切分) | |
|---|---|---|
| Softmax 统计 | 部分列,跨 warp 合并 | 完整行,独立完成 |
| 输出 写入 | 多 warp 更新同一行 | 各 warp 独占不同行 |
| 通信开销 | shared memory / shuffle | 零 |
Causal Mask 优化
因果注意力要求位置 只关注 。V2 在块级别判断:完全在对角线以下的块正常算,完全在对角线以上的块直接跳过,跨对角线的块逐元素 mask。对角线以上约占一半,块级跳过使因果注意力实际开销接近非因果的一半。
LogSumExp 合并统计量
V1 向 HBM 写入 和 两个向量。V2 将它们合并为一个标量 ,存储减半。
为什么一个值就够?回忆 Softmax 的定义:
取对数:
因此 。反向传播需要恢复 时,只需从 HBM 读取 (而非分别读 和 ),配合重计算得到的 ,一步指数即可还原。 本质上就是 log-normalizer:(数值稳定形式为 )。
性能对比6
| 序列长度 | V1 TFLOPS | V2 TFLOPS | V2 利用率 |
|---|---|---|---|
| 1024 | 124 | 196 | 63% |
| 2048 | 136 | 218 | 70% |
| 4096 | 141 | 227 | 73% |
| 8192 | 138 | 222 | 71% |
结语
FlashAttention 的核心原则是:在 GPU 算力远超带宽的硬件现实下,用计算换 IO。 Tiling 把数据锁在 SRAM 里,Online Softmax 在不知道全局统计量时实现精确归一化,重计算替代存储。每一步都在将算子从 memory-bound 推向 compute-bound。
从 V1 到 V3 的迭代展现了系统优化的层级递进:V1 解决算法问题(如何正确分块),V2 解决工程问题(如何利用硬件并行度),V3 解决硬件适配问题(如何压榨特定架构的每一条执行通道)。层次在变,原则不变:找到硬件的真实瓶颈,然后针对性地设计算法。
参考资料
-
Roofline: An Insightful Visual Performance Model for Multicore Architectures (Williams et al., 2009) ↩
-
FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness (Dao et al., 2022) ↩ ↩2 ↩3 ↩4 ↩5 ↩6 ↩7
-
Online Normalizer Calculation for Softmax (Milakov & Gimelshein, 2018) ↩ ↩2
-
FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning (Dao, 2023) ↩ ↩2 ↩3 ↩4
-
FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision (Shah, Dao et al., 2024) ↩ ↩2