FlashAttention 入门:从 FA1 到 FA4 的 IO 感知注意力
本页目录
FlashAttention 入门:从 FA1 到 FA4 的 IO 感知注意力
注意力层是把 Transformer 扩展到长序列时的主要瓶颈:它的时间与内存开销都随序列长度平方增长。FlashAttention 是一族让注意力保持精确(exact)的同时显著提速、省内存的算法,其核心理念是把 GPU 的内存流量——而不只是 FLOPs——当作需要优化的资源。本文用原创的图示与推导讲清核心思想:GPU 内存层级、分块(tiling)、在线 softmax(online softmax)、反向重计算(recomputation),然后是 FA1 → FA2 → FA3 → FA4 的时间线、独立 flash-attn 包与 PyTorch scaled_dot_product_attention 的区别,以及实践中真正要紧的硬件注意事项。
本文有两条规则。第一,所有性能数字均为引用,来自原始论文、Dao-AILab 官方仓库与 PyTorch 官方文档,每个数字都绑定其测量硬件;本文没有任何本地基准结果。第二,本文写作所用的工作站 GPU 是 TITAN RTX(Turing 架构,计算能力 7.5)——这不是巧合,因为官方的 FlashAttention-2/3/4 CUDA 内核无法在 Turing 上运行,最后一节把这个限制转化为一份诚实的验证方案。
文中所有”截至”表述均指 2026-08-09。
注意力为什么是平方复杂度#
以单头为例,设 query、key、value 矩阵为 $Q, K, V \in \mathbb{R}^{N \times d}$($N$ 个 token,头维度 $d$),缩放点积注意力为:
$$S = \frac{Q K^{\top}}{\sqrt{d}}, \qquad P = \operatorname{softmax}(S) \text{(按行)}, \qquad O = P V,$$
其中 $O \in \mathbb{R}^{N \times d}$ 是输出。算术开销是 $O(N^2 d)$ 次 FLOPs:两次 $N \times N \times d$ 的矩阵乘法。这部分是固有的——真正让人头疼的是内存开销。得分矩阵 $S$ 和概率矩阵 $P$ 各有 $N^2$ 个元素,而标准实现会把两者都物化(materialize)在 HBM(GPU 主存)里。
一个纯算术的算例:$N = 8192$ 个 token、头维度 $d = 128$、fp16(2 字节)时,单头仅存放 $S$ 和 $P$ 就需要 $2 \times 8192^2 \times 2$ 字节 ≈ 268 MB。32 个头就是 ≈ 8.6 GB——这还没算模型权重、KV 缓存和激活值。到 $N = 64,K$(长上下文场景),单头需要 ≈ 17 GB,32 个头 ≈ 550 GB,任何现有 GPU 都装不下。这正是稀疏、低秩、哈希等近似注意力方法流行的原因:它们想砍掉 $N^2$ 成本。FlashAttention 走了另一条路——它计算相同的矩阵,但绝不让它们出现在 HBM 里。
GPU 内存层级与 IO 感知#
GPU 有两个相关的内存层级。HBM 容量大(几十 GB)但相对慢:A100 约 1.5 TB/s、H100 约 3.35 TB/s(NVIDIA 数据手册),TITAN RTX 约 672 GB/s。SRAM 是紧挨着计算单元的片上内存:每个流式多处理器(SM)只有几十到几百 KB(FlashAttention v1 的 README 给出 T4 为 64 KB;NVIDIA 数据手册列出 A100 每 block 可用共享内存最高 164 KB、H100 每 SM 为 228 KB)。SRAM 比 HBM 快一个数量级,但装不下 $N \times N$ 矩阵。
图 1. GPU 内存层级与 FlashAttention 的数据流。
+------------------------------------------------------------------+
| HBM —— 主存,几十 GB |
| Q, K, V, O (A100 ~1.5 TB/s,H100 ~3.35 TB/s, |
| TITAN RTX ~672 GB/s) |
| |
| | 按块加载(Q_i, K_j, V_j) ^ 写回 O 分块 |
| v | |
| +----------------------------------------------+ |
| | SRAM —— 片上,每 SM,约 64–228 KB | |
| | S_ij、P_ij、运行中的 m、l、O_i 都放在这里 | |
| | 注意力逐块在本地算完 | |
| +----------------------------------------------+ |
+------------------------------------------------------------------+
FLOPs 数量与标准注意力完全相同;改变的是中间矩阵存放在哪里。FA1 论文的关键观察是:很多运算是内存受限(memory-bound)的——它们的运行时间由内存访问次数决定,而不是由算术量决定。标准注意力的一次前向正是这种情况:
图 2. HBM 流量对比:标准注意力 vs FlashAttention(前向)。
标准注意力: FlashAttention:
Q, K --HBM--> S = QK^T / sqrt(d) Q、K、V 按块流入;
S --HBM--> 逐行求最大/求和(softmax) S_ij 与 P_ij 在 SRAM 内算完
P --HBM--> O = P V 即被消费,从不写回
HBM 流量: Theta(Nd + N^2) HBM 流量:Theta(N^2 d^2 / M)
额外内存: O(N^2) 额外内存:O(N)
标准注意力要把 $S$ 写进 HBM、读回来做 softmax、再写 $P$、再读回来乘 $V$——四次平方级的往返。FlashAttention 的 IO 复杂度分析(2205.14135 的定理 2)证明标准注意力需要 $\Theta(Nd + N^2)$ 次 HBM 访问,而 FlashAttention 只需 $\Theta(N^2 d^2 / M)$,其中 $M$ 是 SRAM 大小。对常见的 $d = 64$–$128$、$M$ 在几十到几百 KB 的量级,$M \gg d^2$,于是 FlashAttention 的 HBM 访问次数少很多倍。论文还给出了匹配的下界:任何精确注意力算法都不可能在所有 SRAM 大小上渐近地优于这一结果。
FlashAttention-1:分块与在线 softmax#
把一切融合进单个内核的障碍是 softmax:$P$ 的第 $i$ 行依赖该行全部 $N$ 个得分的最大值和总和,朴素实现必须看完整行才能输出任何结果。FlashAttention(2205.14135 的算法 1)把两个经典技巧组合起来解决这个问题——分块(tiling)与在线(online)softmax——反向传播再用重计算(下一节)。
分块。把 $Q$ 切成大小为 $B_r \times d$ 的行块 $Q_i$,把 $K, V$ 切成大小为 $B_c \times d$ 的列块 $K_j, V_j$。论文令 $B_c = \lceil M / 4d \rceil$、$B_r = \min(\lceil M / 4d \rceil, d)$,使 $Q_i, K_j, V_j, S_{ij}$(分别是 $B_r \times d$、$B_c \times d$、$B_c \times d$、$B_r \times B_c$)四块能同时放进 SRAM。内核遍历所有 $(i, j)$ 分块,计算 $S_{ij} = Q_i K_j^{\top} / \sqrt{d}$ 并汇入运行中的输出——$N \times N$ 矩阵从头到尾不存在。
图 3. 注意力计算的分块。每个分块 S_ij 在 SRAM 内算完、softmax、消费。
K 分块(每个 B_c x d)->
+----------------------------------+
| S_11 S_12 S_13 ... S_1,nc | 单个分块内部:
| S_21 S_22 S_23 ... S_2,nc | K_j (B_c x d) V_j (B_c x d)
| ... | \ /
| S_nr1 S_nr2 ... S_nr,nc | S_ij = Q_i K_j^T / sqrt(d)
+----------------------------------+ P_ij = softmax(S_ij)
Q 分块(每个 B_r x d) O_i += P_ij V_j
在线 softmax。标准技巧是把 softmax 写成”减去最大值”的形式,只维护运行中的最大值与归一化量。设处理完 key 分块 $1 \dots j$ 后,$m^{(j)}$ 是运行中的行最大值,$\ell^{(j)}$ 是 $\exp(s_k - m^{(j)})$ 的运行和,$O^{(j)}$ 是运行中的未归一化输出 $\sum_{k \le j} \exp(s_k - m^{(j)}) v_k$。不变量是:任何一步的”真输出”都等于 $O^{(j)} / \ell^{(j)}$,最后一步除一次即可。当带得分 $s_{j+1}$ 的新分块到达时,先抬升最大值 $m’ = \max(m^{(j)}, \max s_{j+1})$,再用 $\alpha = e^{m^{(j)} - m’}$ 缩放旧状态、用 $\beta = e^{\max s_{j+1} - m’}$ 缩放新分块:
$$\ell^{(j+1)} = \alpha, \ell^{(j)} + \beta, \ell_{j+1}, \qquad O^{(j+1)} = \alpha, O^{(j)} + \beta, \tilde{P}{j+1} V{j+1},$$
其中 $\tilde{P}{j+1} = \exp(S{j+1} - \max s_{j+1})$,$\ell_{j+1} = \sum \tilde{P}{j+1}$(按行)。为什么这样是对的:旧行的指数项此前都相对旧最大值 $m^{(j)}$ 缩放了 $e^{m^{(j)}}$ 倍;把参考最大值从 $m^{(j)}$ 换成 $m’$,恰好把每个旧指数项乘以 $\alpha$、每个新分块指数项乘以 $\beta$。输出 $O$ 是 value 向量以这些指数为权重的加权和,因此必须做完全相同的缩放。对 $j$ 做简短归纳即可确认不变量:任意前缀之后 $O^{(j)} = \sum{k \le j} \exp(s_k - m^{(j)}) v_k$,故最终的 $O / \ell$ 精确等于 $\operatorname{softmax}(S) V$。
图 4. 新 key 分块到达时的在线 softmax 状态更新。
更新前状态: 新分块 j+1: 更新后:
m = 运行中最大值 m' = max(m, max s_j+1) m <- m'
l = sum exp(s_k - m) a = exp(m - m') l <- a*l + b*l_j+1
O = sum exp(s_k - m) v_k b = exp(max s_j+1 - m') O <- a*O + b*P~ V_j+1
因果掩码几乎免费:因果注意力只需跳过 $j > i$ 的分块(或用 $-\infty$ 掩掉),工作量减半。FA1 论文还把方案扩展到块稀疏(block-sparse)注意力,论文报告它在保留分块内精确性的同时,比已知的近似注意力方法都快。
因此前向只使用线性额外内存——论文定理 1:算法 1 在输入输出之外只用 $O(N)$ 额外内存返回 $\operatorname{softmax}(QK^{\top})V$——且 HBM 访问为 $\Theta(N^2 d^2 / M)$ 而非 $\Theta(Nd + N^2)$。
反向传播:用重计算代替存储#
训练还需要梯度。朴素反向需要再次用到 $P$(或 $S$),存储它要 $O(N^2)$ 内存。FlashAttention 只保存小的逐行统计量——$O$ 与每行的 log-sum-exp $m_i$——然后在反向内核里重计算每个 $S_{ij}$ 分块。设输出梯度为 $dO$,且 $\tilde{P}{ij} = \exp(S{ij} - m_i)$:
$$dV_j = \tilde{P}{ij}^{\top} dO_i, \qquad dS{ij} = \tilde{P}{ij} \odot \left( dO_i V_j^{\top} \right), \qquad dQ_i \mathrel{+}= dS{ij} K_j, \qquad dK_j \mathrel{+}= dS_{ij}^{\top} Q_i.$$
图 5. 反向传播:重计算,而不是存储。
前向只存: O、m(逐行 log-sum-exp) —— 不存 N x N 矩阵
反向: 重新加载 Q、K、V 分块
S_ij = Q_i K_j^T / sqrt(d) (重计算)
P_ij = exp(S_ij - m_i) (重计算)
dV_j = P_ij^T dO_i
dS_ij = P_ij o (dO_i V_j^T)
dQ_i += dS_ij K_j, dK_j += dS_ij^T Q_i
这是用少量额外算术换取大幅减少的 HBM 访问;论文明确指出重计算版反向更快而非更慢,正是因为 HBM 流量主导了运行时间。官方仓库的 flash_attn_func 文档给出了确定性结论:前向永远是确定性的;默认反向则跨线程块用原子加(atomic add)累加 $dQ$/$dK$,逐次运行在比特级不确定。仓库提供 deterministic=True 开关,文档注明”略慢且占用更多内存”。FlashAttention-4 的论文后来也把”减少反向传播中的原子加”列为明确的设计目标之一。
FlashAttention-2:更好的工作划分#
FA2 论文(2307.08691 )从一个令人清醒的测量出发:FA1 只达到 GPU 理论 FLOPs/s 的 25–40%——远低于优化 GEMM 的水平。诊断结果是:线程块与线程束(warp)之间的工作划分欠佳,造成占用率低和不必要的共享内存流量。三个修复直接对应:
- 减少非 matmul 的 FLOPs:用更廉价的数学缩放运行中的统计量(例如能除就不重新求幂),让更多周期花在张量核的 matmul 上。
- 即使单个头也按序列维度并行:FA1 中一个头的输出只能由有限的线程块计算;FA2 让网格覆盖 $(batch, heads, \text{seqlen}_q \text{ 分块})$,长序列因此获得更多线程块(更高占用率)——而且前向不需要原子操作,因为每个输出分块只属于一个线程块。
- 更合理的块内线程束划分:不再按输出行切分线程束(那会强迫共享内存往返),而是把 $K, V$ 的列分给不同线程束,使各线程束的部分结果可以用更少通信合并。
图 6. FA2 前向:覆盖 (batch, heads, seqlen_q 分块) 的线程块网格。
网格:batch x heads x (N / B_r) 个线程块
seqlen_q 分块 ->
+----+----+----+----+----+
| C | C | C | C | C | 每个 C 独占一个 B_r x d 输出分块,
| C | C | C | C | C | 流式地把 K/V 分块送进 SRAM,
+----+----+----+----+----+ 前向无需原子操作
论文摘要报告”相比 FlashAttention 约 2 倍加速,在 A100 上达到理论峰值 FLOPs/s 的 50–73%“,端到端训练 GPT 风格模型达到每张 A100 225 TFLOPs/s(72% 模型 FLOPs 利用率)。官方 README 的更新日志显示这条发布线随后长成了完整的推理工具箱:变长(varlen,非填充)入口、滑窗注意力、ALiBi、确定性反向、分页 KV 缓存、softcapping 与 torch.compile 兼容。硬件方面 README 说得很直白:FA2 CUDA 内核支持 Ampere、Ada 与 Hopper(sm80+);对 Turing 则指向一个独立的社区仓库 ssiu/flash-attention-turing ,”它在 Turing 上支持 FlashAttention 核心功能子集”。
FlashAttention-3:Hopper 异步与 FP8#
FA3(2407.08608 )面向 H100,在 H100 上 FA2 只有约 35% 的利用率。论文把原因归结为 FA2 没有利用 Hopper 的新硬件能力,并提出三项技术:
- warp 专用化 + TMA:让一部分 warp 专职通过张量内存加速器(TMA,异步批量拷贝,HBM → 共享内存)搬运数据,另一部分 warp 专职张量核计算,使内存传输与计算重叠而非交替。
- matmul 与 softmax 交错:把一次迭代中的两次 GEMM 与前一次的 softmax 软件流水化,张量核不必等待指数/softmax 单元。
- 带块量化与非相干处理的 FP8:按块量化得分,量化前先施加随机(非相干)旋转,论文报告这显著降低了 FP8 数值误差;FA3 的 FP8 经验证比基线 FP8 注意力的误差低 2.6 倍。
图 7. FA3 在 Hopper 上的 warp 专用化流水线。
生产者 warp: TMA:HBM -> SRAM (在……的同时取 K_{j+1}, V_{j+1})
消费者 warp: MMA:S_ij = Q_i K_j^T | softmax(S_ij)
MMA:O_i += P_ij V_j | (交错、流水化)
时间 -> 重叠:算第 j 块的同时取第 j+1 块
报告数字(H100,出自摘要):相比 FA2 加速 1.5–2.0 倍,FP16 最高 740 TFLOPs/s(75% 利用率),FP8 接近 1.2 PFLOPs/s。截至 2026-08-09,官方仓库仍把 FA3 标注为测试版(beta):要求 H100/H800、CUDA ≥ 12.3(推荐 12.8),目前只发布 FP16/BF16 前向+反向、FP8 仅前向。
FlashAttention-4:Blackwell 软硬协同设计#
FA4(2603.05451 ,2026 年 3 月)是对硬件变化的回应:在 Blackwell(B200/GB200)上,张量核吞吐大约翻倍,而共享内存带宽与指数单元几乎没有增长——FA3 为 Hopper 调优的流水线因此把新瓶颈晾在一边。FA4 的技术直接映射这个不对称性:
- 重新设计流水线:完全异步的 MMA 操作与更大的分块,让翻倍的张量核保持饱和,不必等待共享内存。
- 软件模拟指数与条件式 softmax 重缩放,把工作从几乎没增长的专用函数单元上挪走。
- 张量内存(tensor memory)与 2-CTA MMA 模式,削减共享内存流量,并——明确地——削减反向传播中的原子加。
论文报告在 B200 上、BF16 下最高比 cuDNN 9.13 快 1.3 倍、比 Triton 基线快 2.7 倍,达到 1613 TFLOPs/s(71% 利用率)。一个值得注意的工程转变:FA4 全部用 CuTeDSL(一种内嵌于 Python 的 DSL)实现,论文报告编译时间比 C++ 模板内核快 20–30 倍;官方 README 的安装方式是 pip install flash-attn-4,面向 Hopper 与 Blackwell(可选 cu13 extra 适配 CUDA 13)。
表 1 汇总时间线。所有数字均出自对应论文的引用,本文未复现。
表 1. FlashAttention 时间线(所有数字均为引用,未复现)。
| 世代 | 论文 / 发布 | 目标硬件 | 核心思想 | 报告性能 |
|---|---|---|---|---|
| FA1 | 2205.14135 ,NeurIPS 2022 | Ampere;v1 发布版也支持 Turing | 分块、在线 softmax、反向重计算、块稀疏扩展 | A100 上比标准注意力快 2–4 倍(序列 128–4K);BERT-large 端到端比 MLPerf 1.1 纪录快 15%;每张 A100 189 TFLOPs/s(60.6% MFU) |
| FA2 | 2307.08691 ,ICLR 2024 | Ampere、Ada、Hopper(sm80+) | 工作划分、更少非 matmul FLOPs、序列维度并行 | 约 2 倍于 FA1;A100 峰值 FLOPs/s 的 50–73%;每张 A100 225 TFLOPs/s(72% MFU) |
| FA3 | 2407.08608 ,2024(beta) | Hopper H100/H800 | TMA、warp 专用化、matmul/softmax 交错、FP8 | H100 上 1.5–2.0 倍于 FA2;FP16 740 TFLOPs/s(75%);FP8 约 1.2 PFLOPs/s |
| FA4 | 2603.05451 ,2026 | Hopper + Blackwell(B200/GB200) | CuTeDSL、异步 MMA 流水线、模拟 exp/重缩放、张量内存、2-CTA MMA | B200 BF16 下 1.3 倍于 cuDNN 9.13、2.7 倍于 Triton;1613 TFLOPs/s(71%) |
独立 flash-attn 与 PyTorch SDPA 的区别#
两种接触 FlashAttention 风格内核的途径完全不同,却经常被混为一谈:
- 独立包(Dao-AILab/flash-attention )是论文背后的内核库。它提供
flash_attn_func、flash_attn_qkvpacked_func、flash_attn_varlen_func、flash_attn_with_kvcache等接口,README 记载了很长的功能清单:dropout、因果与滑窗掩码、ALiBi、MQA/GQA、分页 KV 缓存、融合旋转位置编码,以及deterministic反向开关。安装方式是编译 CUDA 内核(pip install flash-attn --no-build-isolation;需要 CUDA/ROCm 工具链、PyTorch ≥ 2.2、Linux);FA3 与 FA4 以独立包发布(hopper/子目录下的flash-attn-3、flash-attn-4)。 - PyTorch SDPA(
torch.nn.functional.scaled_dot_product_attention)是框架级 API,带自动分派器。PyTorch 2.13 文档 列出三个后端——FlashAttention-2 风格内核、内存高效(memory-efficient,xFormers 风格)内核、C++ math 实现——较新版本还有 cuDNN 后端。后端自动选择,也可以用torch.nn.attention.sdpa_kernel()或torch.backends.cuda.enable_flash_sdp()/enable_mem_efficient_sdp()/enable_math_sdp()强制或禁用。flash 后端是 PyTorch 自己的实现(源码甚至注释着”FlashAttentionV2 requires that head dimension be a multiple of 8”),不是 Dao-AILab 包,且有自己的一套限制:不接受显式attn_mask(只有is_causal标志)、因果 + 非方形序列长度会被拒绝、头维度 ≤ 256。
实践上的区别:SDPA 一次调用就会为你的输入、在任何 GPU 上选一个安全的内核;独立包则给你完整的内核功能集。两者不能当作可互换的 drop-in——文档警告”选择不同后端内核时,本函数的输出可能不同”,而且 math 后端对 fp16/bf16 输入会以 fp32 保存中间量,这就是 SDPA 结果在不同后端、不同 GPU 之间会略有差异的原因。
硬件、数据类型、确定性与填充注意事项#
表 2 汇总了”什么支持什么”,依据是官方 README 与 PyTorch v2.13 的分派代码(sdp_utils.cpp ,其中把 flash 后端限制在 sm80–sm121,把内存高效后端限制在 sm50–sm121)。
表 2. FlashAttention 家族内核的支持矩阵(截至 2026-08-09)。
| 实现 | 最低 GPU(CC) | fp16 | bf16 | fp8 | 头维度 | 备注 |
|---|---|---|---|---|---|---|
| FA1 v1(官方) | Turing sm75+ | 支持 | 支持,sm80+ | 不支持 | ≤ 128,8 的倍数;反向 > 64 需 A100/H100 | varlen API;128/256 token 分块 |
| FA2(官方 CUDA) | Ampere sm80+ | 支持 | 支持 | 不支持 | 最高 256 | Turing 用户被指向社区 Turing 仓库 |
| FA3(官方,beta) | Hopper H100/H800 | 支持 | 支持 | 仅前向 | 视内核而定 | CUDA ≥ 12.3,推荐 12.8 |
| FA4(官方) | Hopper + Blackwell | — | 支持(已基准) | — | 视内核而定 | pip install flash-attn-4;CuTeDSL |
| PyTorch SDPA flash | sm80–sm121 | 支持 | 支持 | sm90+,若启用 FA3 | ≤ 256 | 无 attn_mask;拒绝非方形因果 |
| PyTorch SDPA 内存高效 | sm50–sm121 | 支持 | sm80+ | 不支持 | 对齐 8(fp16)/ 4(fp32) | sm80 以下仅限 fp16/fp32 |
| PyTorch SDPA math | 任意(含 CPU) | 支持 | 支持 | 不支持 | 任意 | fp32 累积;支持 fp64 |
| ssiu/flash-attention-turing | sm75 | 支持 | 见仓库 | 不支持 | 核心子集 | 官方 README 点名的社区移植 |
实践中会咬人的注意事项:
- 数据类型。融合内核是 fp16/bf16 的天下。bf16 需要 sm80+(FA1 README 与 PyTorch 的数据类型门槛一致);sm80 以下,PyTorch 的融合内核只接受 fp16。fp32 只能走 math 或内存高效后端。FP8 是 FA3 时代的领域(e4m3),且官方 beta 只有前向;FA4 的头条数字是 BF16。
- 确定性。独立包 FA 前向永远确定;默认反向在比特级不确定(原子累加),
deterministic=True是显式开关。PyTorch 文档提到输出随后端而异,cuDNN 路径可能选择非确定性算法。需要可复现运行时,就强制 math 后端或使用确定性开关——代价是在你的负载上实测出来的性能损失。 - 填充与形状。融合内核要求头维度是 8 的倍数(FA1 直接断言;PyTorch 在 composite 层把头维度补齐到 8 的倍数)。FA1 把头维度上限设为 128(非 A100/H100 上反向为 64);FA2 与 PyTorch flash 后端允许到 256。PyTorch flash 后端拒绝因果 +
seqlen_q != seqlen_k的组合。独立包通过flash_attn_varlen_func(以及 FA1 时代的cu_seqlensAPI)支持变长序列、无需填充,FA4 的头条特性就是明确的非填充(un-padded)注意力;cuDNN SDPA 后端历史上要求seq_kv是 64 的倍数(cuDNN 8.9.6 之前)、头维度是 8 的倍数——于是某个后端能跑的填充形状,换一个后端可能悄悄回退。
本机(Turing)能跑什么,以及我们提议的诚实验证方案#
本文工作站的 GPU 是 TITAN RTX(TU102,计算能力 7.5,24 GB HBM2,按 NVIDIA 规格约 672 GB/s)。对照表 2 折叠到本机:
- 能跑:官方 FA1 v1 内核(fp16;头维度 ≤ 128,反向 ≤ 64)——FA1 是唯一在官方 CUDA 内核清单里列出 Turing 支持的世代;PyTorch SDPA 的内存高效(fp16/fp32)与 math 后端;社区
ssiu/flash-attention-turing移植版。 - 不能跑:官方 FA2、FA3、FA4 CUDA 内核——FA2 需要 Ampere 或更新(sm80+),FA3 需要 H100/H800,FA4 需要 Hopper/Blackwell。PyTorch 的 SDPA flash 后端同样被限制在 sm80–sm121,在本 GPU 上不会被选中;bf16 融合内核也不可用(bf16 需要 sm80+)。
与其报告任何我们没有实际跑过的东西,不如给出四条面向 Turing 的提议验证方案,以及依据所引来源它们各自应当产生的结果。本文写作时没有执行其中任何一条;请把它们当作具体计划,而不是结果。
- 后端探测(PyTorch)。在本 GPU 上用
sdpa_kernel([SDPBackend.FLASH_ATTENTION])调用F.scaled_dot_product_attention。预期:分派器拒绝,并给出 v2.13 源码中的警告原文——”Flash attention only supports gpu architectures in the range [sm80, sm121]“——然后改用[SDPBackend.EFFICIENT_ATTENTION, SDPBackend.MATH]重跑,验证输出与 math 后端在 fp16 容差内一致。 - 算法参照(eager PyTorch)。用纯 torch 算子实现算法 1(分块 + 在线 softmax + 最终重缩放,即上图 4 的流程),fp16、头维度 64/128、序列 512/2048、含与不含因果,与同一张量上的 SDPA math 后端对比;沿用官方仓库的测试准则:最大误差不超过基线实现误差的约 2 倍。这条可以在 Turing 上、不借助任何融合内核验证本文的推导。
- Turing 上的 Triton(固定旧版本)。改编 OpenAI 的 06-fused-attention.py 教程内核——官方 FA README 把它当作可读性最好的参考实现——搭配 Triton 2.1.x 发布版在 TITAN RTX 上运行;该版本 README 写明”NVIDIA GPUs (Compute Capability 7.0+)“。注意这个带日期的坑:当前的 Triton(main,2026 年)要求计算能力 8.0+,所以这条必须用旧版本;预期结果是 fp16 内核可运行、输出与 math 后端一致。
- 社区 Turing 内核。安装
ssiu/flash-attention-turing(官方 FA2 README 为 Turing 点名的仓库),先在 TITAN RTX 上验证其前向/反向与 SDPA 内存高效后端一致,然后才比较计时。
如果你真的在 Turing GPU 上跑这些,最有意思的问题正是论文 IO 分析预测的那几个:在 SRAM 只有 64 KB 级、没有异步拷贝硬件的 GPU 上,加速比还剩多少(FA1 的 T4 实测提示前向 2.5–4.5 倍是合理区间);以及同等精度下,内存高效的 cutlass 内核与分块在线 softmax 公式相比表现如何。
总结#
- 注意力的平方级内存开销(而非 FLOPs)才是最初的瓶颈;FlashAttention 通过 IO 感知保持注意力精确:分块 $Q, K, V$,在 SRAM 里算 $S$ 和 $P$,绝不物化 $N \times N$ 矩阵(HBM 流量 $\Theta(N^2 d^2 / M)$ 对比 $\Theta(Nd + N^2)$,见 FA1 论文定理 2)。
- 在线 softmax(运行最大值、运行和、重缩放)让融合成为可能;重计算让反向省内存,并且尽管多了 FLOPs 反而更快。
- FA2 修好了工作划分(约 2 倍于 FA1;A100 峰值的 50–73%);FA3 用足 Hopper 的 TMA/warp 专用化/FP8(H100 上 FP16 740 TFLOPs/s、FP8 约 1.2 PFLOPs/s);FA4 针对 Blackwell 的不对称扩展重新协同设计(B200 上 BF16 1613 TFLOPs/s、1.3 倍于 cuDNN 9.13,CuTeDSL 实现)。
- 独立
flash-attn是内核库;PyTorch SDPA 是自动分派的 API,自带 FA2 风格后端且限制更严。两者不是 drop-in 等价物。 - 硬件现实:官方 FA2/FA3/FA4 内核需要 Ampere/Hopper/Blackwell。在本工作站的 TITAN RTX(Turing,sm75)上,FA1 v1 与 PyTorch 的内存高效/math 后端是诚实的选择;上面四条验证方案是提议、尚未执行——本文不声称任何本地基准数字。
参考资料#
- Tri Dao, Daniel Y. Fu, Stefano Ermon, Atri Rudra, Christopher Ré. FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness . arXiv:2205.14135,NeurIPS 2022。
- Tri Dao. FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning . arXiv:2307.08691,ICLR 2024。
- Jay Shah, Ganesh Bikshandi, Ying Zhang, Vijay Thakkar, Pradeep Ramani, Tri Dao. FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision . arXiv:2407.08608,2024。
- Ted Zadouri, Markus Hoehnerbach, Jay Shah, Timmy Liu, Vijay Thakkar, Tri Dao. FlashAttention-4: Algorithm and Kernel Pipelining Co-Design for Asymmetric Hardware Scaling . arXiv:2603.05451,2026。
- Dao-AILab. flash-attention:FlashAttention 与 FlashAttention-2 官方实现(含 FA3/FA4 发布) . GitHub 仓库(本文引用其 README 与 v1.0.9 tag)。
- PyTorch. torch.nn.functional.scaled_dot_product_attention 与 torch.nn.attention.sdpa_kernel . PyTorch 2.13 文档。
- PyTorch. aten/src/ATen/native/transformers/cuda/sdp_utils.cpp(v2.13.0) . 本文引用的后端能力门槛。
- Shengqi Chen(ssiu). flash-attention-turing:面向 Turing GPU 的 FlashAttention . 官方 FA2 README 点名的社区仓库。
- Triton. triton-lang/triton README 兼容性说明 (main:NVIDIA 计算能力 8.0+;v2.1.0 tag:7.0+)与 06-fused-attention.py 教程 。