FlashAttention 完全拆解(一):IO 感知与分块注意力算法

6646 字
33 分钟
FlashAttention 完全拆解(一):IO 感知与分块注意力算法

背景:标准注意力实现的三个瓶颈#

FlashAttention 于 2022 年由斯坦福大学的 Tri Dao 等人提出,发表在 NeurIPS 2022 上。它是过去几年对 Transformer 训练和推理影响最大的系统优化之一:这篇论文之前,注意力(attention)实现的速度和显存占用长期受制于一个简单的工程事实——中间矩阵太大,必须反复读写慢速的显存。本站的 FlashAttention 完全拆解(四) 已经讨论过这条技术路线的最终形态,本文回到起点,把 FlashAttention-1/2 这代经典设计彻底讲透。

注意力层的计算是:

Attention(Q,K,V)=softmax(QKTd)V\text{Attention}(Q, K, V) = \text{softmax}\left(\frac{QK^T}{\sqrt{d}}\right)V

其中 Q,K,VRN×dQ, K, V \in \mathbb{R}^{N \times d} 分别是查询(query)、键(key)、值(value)矩阵,NN 是序列长度,dd 是每个头的维度(head dimension)。以 dd 为头的维数,QKTQK^T 得到一个 N×NN \times N 的注意力分数矩阵 SS,softmax 按行归一化得到概率矩阵 PP,再乘上 VV 得到输出 OO

在 FlashAttention 出现之前,GPU 上的标准实现(论文里称为 Algorithm 0)把这三步拆成三个独立的 kernel,每个 kernel 都从高带宽内存(High Bandwidth Memory,HBM)读入、算完、写回:

  1. 加载 QQKK,计算 S=QKTS = QK^T,把 N×NN \times NSS 写回 HBM;
  2. 重新读入 SS,计算 P=softmax(S)P = \text{softmax}(S)(含行内最大值归约、指数、求和、归一化),把 PP 写回 HBM;
  3. 重新读入 PPVV,计算 O=PVO = PV,写回 HBM。

这个实现有三个互相纠缠的瓶颈。

瓶颈一:N×NN \times N 中间矩阵的物化(materialization)带来 O(N2)O(N^2) 显存占用。 序列长度每翻一倍,SSPP 的体积翻四倍。序列长度 4096、头维度 128 时,每个矩阵在 FP16 下约 33.5 MB,SSPP 合起来 67 MB;到 16384 时每个约 537 MB;到 65536 时每个约 8.6 GB,两个矩阵超过 17 GB——一张 A100 的显存就被注意力中间结果占掉一大半。这直接限制了模型能处理的上下文长度,也解释了为什么 FlashAttention 之前的主流模型都把序列长度钉死在 2K 附近。

瓶颈二:HBM 访问次数是 Θ(Nd+N2)\Theta(Nd + N^2) 标准的三个 kernel 把 SSPP 各写一次、各读一次,加上 QQKKVVOO 的读写,总共有 Θ(N2)\Theta(N^2) 量级的数据在 HBM 上进进出出。注意这里说的是访问的元素个数,单位量级是 Θ(N2)\Theta(N^2) 而不是 Θ(N)\Theta(N)——这是理解 FlashAttention 全部设计的关键数字。

瓶颈三:标准注意力是内存受限(memory-bound)的,不是计算受限(compute-bound)的。 GPU 的访存和计算资源存在巨大的比例失衡。以 A100 为例:它有 40–80 GB 的 HBM,带宽 1.5–2.0 TB/s;同时每个流式多处理器(streaming multiprocessor,SM)有 192 KB 片上 SRAM,108 个 SM 合计约 20 MB,总带宽估计约 19 TB/s——片上 SRAM 比 HBM 快一个数量级,但容量小三个数量级以上。一个操作是计算受限还是内存受限,通常用算术强度(arithmetic intensity,每字节内存访问对应的算术操作数)来衡量:算术强度低于硬件拐点(ridge point,即峰值算力除以内存带宽)时,时间由访存决定,就是内存受限。

算一个具体数字。N=4096N = 4096d=128d = 128、FP16 精度下,标准注意力完成一次前向约需要 4N2d8.64N^2d \approx 8.6 GFLOP(两次矩阵乘各 2N2d2N^2d,softmax 约 3N23N^2),而 HBM 流量约 134 MB(SSPP 各写读一遍占了大头),算术强度约 64 FLOP/byte。A100 的 FP16 张量核心(Tensor Core)峰值约 312 TFLOPS,除以 2 TB/s 的 HBM 带宽,拐点约 156 FLOP/byte。64 远小于 156,所以标准注意力跑在内存带宽的天花板上,张量核心大部分时间在空等数据——实际吞吐远低于峰值。

在此之前不是没有优化尝试:kernel 融合(kernel fusion)把掩码(masking)、dropout 等逐元素操作合并进 softmax 的 kernel,减少几趟 HBM 往返;cuDNN、xformers 等库对每个 kernel 做精细调优。但所有这些都是在三个 kernel 的框架内打补丁——N×NN \times NSSPP 依然要物化、要写读,Θ(N2)\Theta(N^2) 的 HBM 访问依然存在。FlashAttention 的贡献不是把某个 kernel 调快,而是从根本上改变数据流,让 Θ(N2)\Theta(N^2) 这个量级消失

核心思想:IO 感知(IO-aware)#

FlashAttention 的核心思想用一句话概括:设计注意力算法时,把 HBM 访问次数当作首要优化目标,而不是 FLOPs。

论文用”IO 感知”(IO-aware)来命名这个原则,背后是一条硬件趋势:计算速度的提升长期快于内存速度的提升,因此大量操作从计算受限滑向内存受限。对注意力而言,标准实现的 HBM 访问是 Θ(Nd+N2)\Theta(Nd + N^2),FlashAttention 把它压到 Θ(N2d2/M)\Theta(N^2d^2/M),其中 MM 是片上 SRAM 的大小——典型配置下这比 N2N^2 少了一个数量级以上的常数因子。更关键的是,论文证明了这是精确注意力能到达的下界(lower bound,详见后文)。

实现 IO 感知需要三件套,缺一不可:

  1. 分块(tiling):把 QQKKVV 切成小块,让每个块的处理在片上 SRAM 内完成,N×NN \times N 的中间矩阵从头到尾不落 HBM;
  2. 在线 softmax(online softmax):softmax 需要全局的统计量(行最大值与归一化和),分块后只能看到局部,必须用流式算法在遍历过程中维护并修正这些统计量;
  3. 重计算(recomputation):反向传播需要 SSPP 这两个 N×NN \times N 的中间矩阵,FlashAttention 不存储它们,而是在反向时从块和统计量中重新算出来。

要强调的一点是:FlashAttention 是精确算法(exact algorithm),不是近似——它和标准注意力在数学上完全等价,唯一的差异是浮点运算顺序不同带来的舍入误差。这与当时流行的稀疏注意力(sparse attention)、线性注意力(linear attention)等近似方法有本质区别,也是它能成为事实标准的最重要原因:不需要任何精度权衡,就能白拿加速和显存节省。

机制一:分块(Tiling)与块大小设计#

分块的目标是让”一个 QQ 块 + 一个 KK 块 + 一个 VV 块 + 对应的分数块 SijS_{ij}“能同时放进片上 SRAM。设 SRAM 大小为 MM(A100 上约 192 KB),则块大小取:

Bc=M4d,Br=min(M4d,d)B_c = \left\lceil \frac{M}{4d} \right\rceil, \qquad B_r = \min\left(\left\lceil \frac{M}{4d} \right\rceil, d\right)

其中 BcB_cKKVV 块的”行数”(每个块是 Bc×dB_c \times d),BrB_rQQ 块的”行数”(每个块是 Br×dB_r \times d)。为什么是除以 4?因为 SRAM 里要同时驻留四样东西:QiQ_iBr×dB_r \times d)、KjK_jBc×dB_c \times d)、VjV_jBc×dB_c \times d)、以及中间结果 Sij=QiKjTS_{ij} = Q_iK_j^TBr×BcB_r \times B_c),四份都占 Θ(M)\Theta(M) 的量级,所以每份只能拿到 M/4M/4。三个约束条件分别是:

Bcd=O(M),Brd=O(M),BrBc=O(M)B_c \cdot d = O(M), \quad B_r \cdot d = O(M), \quad B_r \cdot B_c = O(M)

前两个保证 KjK_jVjV_jQiQ_i 装得下,第三个保证它们相乘出来的分数块 SijS_{ij} 也装得下。联立解出来就是上面的取值:BcB_cΘ(M/d)\Theta(M/d)BrB_rmin(Θ(M/d),d)\min(\Theta(M/d), d)——注意 BrB_r 还受 dd 本身限制,因为 Br×BcB_r \times B_c 的约束要求 Br=O(M/Bc)=O(d)B_r = O(M/B_c) = O(d)

代入 A100 的实际数字看看(以 FP16 元素计:192 KB ÷ 2 字节/元素 = 98304 个元素,d=128d = 128):M/(4d)=98304/512=192M/(4d) = 98304/512 = 192,所以 Bc=192B_c = 192Br=min(192,128)=128B_r = \min(192, 128) = 128。也就是说,QQ 块取 128×128128 \times 128KKVV 块取 192×128192 \times 128,分数块是 128×192128 \times 192——四个块合计 128×128+192×128+192×128+128×192=90112128 \times 128 + 192 \times 128 + 192 \times 128 + 128 \times 192 = 90112 个元素,约 180 KB,刚好塞进 192 KB 的 SRAM,还留出一点余量给掩码和 dropout 状态。

这里要说明一个工程细节:上面的公式是论文给出的理论取值,实际 kernel 里块大小还要扣除掩码、dropout 状态、寄存器等开销,而论文的约束 O(M)O(M) 是渐近意义下的,常数因子由实现取舍。因此实际 FlashAttention kernel 常用 Br=Bc=64B_r = B_c = 64128128 这样的取值,论文在块大小选择上的态度是”在寄存器溢出和 SRAM 容量之间手工调优”。理论公式的意义在于揭示块大小只依赖 MMdd,不依赖序列长度 NN——这保证了不管序列多长,每个块的片上占用都是常数,IO 复杂度才能只由 MM 决定。这个推导逻辑比具体数字更重要。

分块之后,QQ 被分成 Tr=N/BrT_r = \lceil N/B_r \rceil 个块,KKVV 各被分成 Tc=N/BcT_c = \lceil N/B_c \rceil 个块。FlashAttention-1 的循环组织是:外层循环遍历 KKVV 块,内层循环遍历 QQ——每加载一个 KjK_jVjV_j 块进 SRAM,就把它和所有 QQ 块各算一遍(论文图 1 的红色箭头是外层,蓝色箭头是内层):

FlashAttention-1 的分块示意(论文图 1 左):外层循环(红色箭头)遍历 K、V 块,内层循环(蓝色箭头)遍历 Q 块,虚线框中的 N×N 注意力矩阵从未在 HBM 中物化(来源:arXiv:2205.14135 图 1)
FlashAttention-1 的分块示意(论文图 1 左):外层循环(红色箭头)遍历 K、V 块,内层循环(蓝色箭头)遍历 Q 块,虚线框中的 N×N 注意力矩阵从未在 HBM 中物化(来源:arXiv:2205.14135 图 1)

图中虚线框代表那个 N×NN \times N 的注意力矩阵——在 FlashAttention 里它从头到尾不存在于任何地方SijS_{ij} 算出来后在 SRAM 里直接参与 softmax 和加权求和,用完后被下一个 Si,j+1S_{i,j+1} 覆盖。同一张图的右边是作者在 GPT-2 上测的加速数据:仅注意力部分就比 PyTorch 标准实现快 7.6 倍。

机制二:在线 softmax——分块的数学支柱#

分块带来一个直接的数学障碍:softmax 是全局操作,而分块之后每个 SijS_{ij} 只包含部分分数。

回忆 softmax 的定义:PP 的第 ii 行是所有 eSi,je^{S_{i,j}} 除以该行的和 jeSi,j\sum_j e^{S_{i,j}}。这个”行和”需要看完整行才能算出来;为了数值稳定,还要先减去行最大值 mi=maxjSi,jm_i = \max_j S_{i,j},否则指数会溢出。标准实现要两趟(先求最大值、再求指数和),而分块流式计算要求一趟遍历完成所有工作——处理某个 KKVV 块时,后面的块还没看到,无法知道这一行的最大值和总和。

解决办法是维护两个随遍历推进的累计量:运行最大值(running max)mim_i运行归一化常数 i\ell_i。处理第 jj 个块时:

minew=max(mi,m~ij),m~ij=rowmax(Sij)m_i^{\text{new}} = \max(m_i, \tilde{m}_{ij}), \qquad \tilde{m}_{ij} = \text{rowmax}(S_{ij})inew=emiminewi+em~ijminew~ij,~ij=rowsum(eSijm~ij)\ell_i^{\text{new}} = e^{m_i - m_i^{\text{new}}} \ell_i + e^{\tilde{m}_{ij} - m_i^{\text{new}}} \tilde{\ell}_{ij}, \qquad \tilde{\ell}_{ij} = \text{rowsum}\left(e^{S_{ij} - \tilde{m}_{ij}}\right)

其中 mim_ii\ell_i 是前 j1j-1 个块累积下来的统计量(以”当时的最大值”为基准),m~ij\tilde{m}_{ij}~ij\tilde{\ell}_{ij} 是当前块 SijS_{ij} 自己的行最大值和”以本块最大值为基准”的指数和。注意 i\ell_i 的定义细节:它是以当前最大值 mim_i 为基准的指数和——这是整个算法能保持精确的关键。当新块的最大值 m~ij\tilde{m}_{ij} 超过旧值 mim_i 时,之前累积的所有项都要乘上 emiminewe^{m_i - m_i^{\text{new}}} 这个缩放因子(rescaling factor),把基准从旧最大值切换到新最大值;而当最大值没变时,缩放因子是 e0=1e^0 = 1,退化为普通累加。

输出累加器 OiO_i 的更新同样带重缩放:

Oidiag(inew)1(diag(i)emiminewOi+em~ijminewP~ijVj)O_i \leftarrow \text{diag}(\ell_i^{\text{new}})^{-1}\left(\text{diag}(\ell_i)\, e^{m_i - m_i^{\text{new}}} O_i + e^{\tilde{m}_{ij} - m_i^{\text{new}}} \tilde{P}_{ij} V_j\right)

其中 P~ij=exp(Sijm~ij)\tilde{P}_{ij} = \exp(S_{ij} - \tilde{m}_{ij}) 是本块的”未归一化概率”块。所有块处理完后,做最后一次归一化 OiOi/iO_i \gets O_i / \ell_i,得到的就是精确的 softmax(QiKT)V\text{softmax}(Q_iK^T)V

为什么这套递归在数学上精确?可以用归纳不变式(inductive invariant)理解:每一步结束时,不变式”Oi/iO_i / \ell_i = 前 jj 个块对应的精确 softmax 加权和”始终成立。当最大值不更新时只是普通累加;当最大值增大时,旧项和新项被同一个标量重缩放——同一行的所有项始终除以同一个 i\ell_i,比例关系不变,所以最终结果与”先看完全部行再归一化”完全一致。算法做的只是把除法推迟到了最后一步。

这套机制在本站的《浮点数与数值稳定性》一文中有完整的推导和数值分析(包括为什么它不会溢出、为什么下溢只发生在”本来就该被忽略”的项上),这里不再重复。值得单独强调的直觉是:在线 softmax 让所有中间值始终保持在 [0,1][0, 1] 量级——这正是它数值稳定的原因,也是分块注意力能存在的数学前提。后面 FlashAttention-4 的”条件重缩放”、DeepSeek 的 FP8 注意力等工程优化,都是在这套机制上做加减法。

完整前向算法#

把三件套合起来,就是论文的 Algorithm 1(FlashAttention 前向传播):

输入:Q, K, V ∈ R^{N×d} 位于 HBM,片上 SRAM 大小为 M
1: 块大小 B_c = ⌈M/(4d)⌉,B_r = min(⌈M/(4d)⌉, d)
2: 在 HBM 初始化 O = (0)_{N×d},ℓ = (0)_N,m = (-∞)_N
3: 把 Q 分成 T_r = ⌈N/B_r⌉ 个块;把 K、V 各分成 T_c = ⌈N/B_c⌉ 个块
4: 把 O、ℓ、m 分成对应的 T_r 个块
5: for j = 1 to T_c do
6: 加载 K_j、V_j 到 SRAM
7: for i = 1 to T_r do
8: 加载 Q_i、O_i、ℓ_i、m_i 到 SRAM
9: 在片上计算 S_ij = Q_i K_j^T ∈ R^{B_r×B_c}
10: 在片上计算 m̃_ij = rowmax(S_ij),P̃_ij = exp(S_ij − m̃_ij),ℓ̃_ij = rowsum(P̃_ij)
11: 在片上计算 m_i^new = max(m_i, m̃_ij),ℓ_i^new = e^{m_i−m_i^new}·ℓ_i + e^{m̃_ij−m_i^new}·ℓ̃_ij
12: 更新 O_i ← diag(ℓ_i^new)^{-1}(diag(ℓ_i)·e^{m_i−m_i^new}·O_i + e^{m̃_ij−m_i^new}·P̃_ij·V_j),写回 HBM
13: 把 ℓ_i、m_i 更新为 ℓ_i^new、m_i^new,写回 HBM
14: end for
15: end for
16: 返回 O

逐行看这个算法的关键决策:

  • 第 1 行的块大小只依赖 MMdd,前面已经推导过;
  • 第 9 行是两处矩阵乘中的第一处Sij=QiKjTS_{ij} = Q_iK_j^T 在张量核心上完成,这是注意力计算中效率最高的部分;
  • 第 10–11 行是在线 softmax 的状态更新,全部在 SRAM 内完成,SijS_{ij} 算完即被消费;
  • 第 12 行是第二次矩阵乘P~ijVj\tilde{P}_{ij}V_j,把概率块和值块相乘累加进输出。注意这里 OiO_i每次内层迭代都写回 HBM——这是论文表述层面的简化,实际 kernel 里 OiO_i 留在 SRAM/寄存器中,整行 ii 处理完才写回一次;
  • 第 13 行mim_ii\ell_i 这两个 BrB_r 维的小向量是唯一需要跨块传递的状态——它们记录了”到目前为止看了多少”,是重计算(反向传播)时需要保存的全部统计信息。

整体复杂度(论文定理 1):算法使用 O(N2d)O(N^2d) 次浮点运算(和标准实现同量级),除输入输出外只需要 O(N)O(N) 的额外内存——mm\ellNN 个标量。显存从 O(N2)O(N^2) 降到 O(N)O(N),这就是”线性注意力内存”的由来,且不牺牲任何精度

IO 复杂度分析:从 Θ(Nd+N2)\Theta(Nd + N^2)Θ(N2d2/M)\Theta(N^2d^2/M)#

论文把 HBM 访问量的分析形式化为 IO 复杂度(IO complexity),这是理解 FlashAttention 为什么快、以及为什么”已经到极限”的关键。

标准注意力的 HBM 访问是 Θ(Nd+N2)\Theta(Nd + N^2) 逐 kernel 数:S=QKTS = QK^T 一步读 QQKKNdNd 量级)写 SSN2N^2 量级);softmax 一步读 SSPP(各 N2N^2);O=PVO = PV 一步读 PPVVOON2+NdN^2 + Nd)。三项相加,主导项是 Θ(N2)\Theta(N^2)——序列长度平方增长,和显存占用一样。

FlashAttention 的 HBM 访问是 Θ(N2d2/M)\Theta(N^2d^2/M) 逐块数:每个 KKVV 元素从 HBM 只被加载一次(第 6 行,外层循环遍历时加载,之后全部复用);而 QQOO 每个外层块都要被完整读一遍,即总共被读 Tc=N/BcT_c = N/B_c 次,每趟读 O(Nd)O(Nd) 个元素,合计 O(NdTc)=O(N2d/Bc)=O(N2d2/M)O(Nd \cdot T_c) = O(N^2d/B_c) = O(N^2d^2/M)(代入 Bc=Θ(M/d)B_c = \Theta(M/d))。KKVV 各一次读入是 Θ(Nd)\Theta(Nd),被 N2d2/MN^2d^2/M 主导。

两个量级对比一下:标准是 N2N^2,FlashAttention 是 N2d2/MN^2 d^2/M。在典型配置(d=64d = 64128128M100M \approx 100 KB 量级)下,d2Md^2 \ll M,所以 FlashAttention 的 HBM 访问比标准实现少一个很大的常数因子——论文在 A100 上实测最多减少约 9 倍。这个因子不随 NN 增长而改变,所以长序列下收益是持续的:无论是 2K 还是 64K 上下文,HBM 流量都从”随 N2N^2 爆炸”变成”随 N2N^2 增长但系数小得多”。

更重要的是下界:精确注意力的 IO 复杂度已经渐近最优。 论文的命题 3(Proposition 3)证明:对 M[d,Nd]M \in [d, Nd] 范围内的任意 SRAM 大小,不存在任何精确注意力算法能以 o(N2d2/M)o(N^2d^2/M) 的 HBM 访问量完成计算。证明思路是取极端情形 M=Θ(Nd)M = \Theta(Nd)(SRAM 大到能装下整个 KKVV),此时 N2d2/M=Θ(Nd)N^2d^2/M = \Theta(Nd),而任何算法至少要把 QQKKVVOO 读入写出一次——这是物理下限,任何算法都逃不掉。

这个下界有深远的实践含义:FlashAttention-2/3/4 的改进空间不在 IO 复杂度本身(那已经是下界),而在常数因子、并行度、硬件利用率。后文会看到,FlashAttention-2 的加速正是这样来的——算法没变,变的是线程块和 warp 怎么分工。

反向传播:用重计算(recomputation)换显存#

前向的分块解决了 SSPP 的物化问题,但反向传播(backward pass)还藏着一个 O(N2)O(N^2) 的坑:计算梯度需要 SSPP。具体来说,设 dOdO 是输出的梯度,标准反演需要:

dV=PTdO,dP=dOVT,dS=P(dPD),dQ=dSK,dK=dSTQdV = P^T dO, \quad dP = dO \cdot V^T, \quad dS = P \circ (dP - D), \quad dQ = dS \cdot K, \quad dK = dS^T \cdot Q

其中 Di=rowsum(dOiOi)D_i = \text{rowsum}(dO_i \circ O_i) 是 softmax 反向的修正项,\circ 表示逐元素乘。如果像标准实现那样把 SSPP 存下来给反向用,显存又回到 O(N2)O(N^2),前向省下的空间全部白费。

FlashAttention 的选择是不存 SSPP,反向时重算。前向结束时,每个 QQ 块只需把两样东西留给反向:输出块 OiO_i,以及统计量 mim_ii\ell_i(合计 O(N)O(N) 个标量)。反向的每个内层迭代从 HBM 加载 QiQ_iKjK_jVjV_jOiO_idOidO_i 到 SRAM,在片上重算 SijS_{ij} 和概率块:

Pij=diag(i)1exp(Sijmi)P_{ij} = \text{diag}(\ell_i)^{-1} \exp(S_{ij} - m_i)

然后用上面的五个公式在块内完成 dVjdV_jdKjdK_jdQidQ_i 的梯度累积。重算的成本只是多一次 QiKjTQ_iK_j^T 的矩阵乘——而这本来就是张量核心最擅长的操作,且在 SRAM 内完成,不产生额外 HBM 流量。

论文把这种设计归类为选择性梯度检查点(selective gradient checkpointing):标准的梯度检查点技术用重算换内存,代价是训练变慢;FlashAttention 的反向虽然 FLOPs 更多(约翻倍),但 HBM 访问从标准反向的 Θ(Nd+N2)\Theta(Nd + N^2) 降到和正向一样的 Θ(N2d2/M)\Theta(N^2d^2/M)——内存少了、FLOPs 多了、反而更快,因为反向同样是内存受限的。论文定理 5 给出了与正向对称的复杂度结论。这大概是”以算换访”在系统优化史上最成功的一次应用。

块稀疏扩展(Block-Sparse FlashAttention)#

论文还顺带把 FlashAttention 扩展成了块稀疏版本(block-sparse FlashAttention):给定一个块级掩码 M~{0,1}N×N\tilde{M} \in \{0,1\}^{N \times N},跳过被掩码盖住的 SijS_{ij} 块的计算,IO 复杂度按稀疏度比例降低——比如稀疏度为 1/s1/s 时 HBM 访问约减少 ss 倍。

这个扩展的学术意义在于:它证明了 FlashAttention 可以作为一种通用原语,让精确注意力和各种近似注意力(稀疏、低秩、核方法)在同一套 IO 框架下比较和实现。本站在《BLASST》里拆解过的动态块稀疏注意力,本质上就是这条线的延续——把”掩码从哪里来”从静态变成运行时决策。不过 FlashAttention 的块稀疏版本由于掩码固定的限制,在真实场景中逐渐被更灵活的方案取代,这里就不展开细讲了。

性能表现#

论文在多个 GPU 和多个模型上做了评测,几组关键数据:

HBM 访问是运行时间的决定因素。 论文图 2 用 GPT-2 medium(序列长度 1024、头维度 64、16 头、batch 64)在 A100 上做微基准:左图对比标准注意力和 FlashAttention 的前向+反向运行时间,中图展示 FlashAttention 前向时间随 HBM 访问减少而下降——访问量越少越快,直到某个拐点(约 9 倍访问量差异):

FlashAttention 微基准(论文图 2):左图为 GPT-2 medium 在 A100 上标准注意力与 FlashAttention 的前向+反向运行时间,HBM 访问是主导因素;中图为前向时间随 HBM 访问减少而下降;右图为块稀疏 FlashAttention 在序列长度 4K 下按稀疏度比例加速(来源:arXiv:2205.14135 图 2)
FlashAttention 微基准(论文图 2):左图为 GPT-2 medium 在 A100 上标准注意力与 FlashAttention 的前向+反向运行时间,HBM 访问是主导因素;中图为前向时间随 HBM 访问减少而下降;右图为块稀疏 FlashAttention 在序列长度 4K 下按稀疏度比例加速(来源:arXiv:2205.14135 图 2)

运行时间和显存占用双双下降。 论文图 3 左侧是前向+反向的端到端运行时间对比,右侧是显存占用——标准注意力(红色线)显存随序列长度平方增长,FlashAttention(蓝色线)线性增长,序列越长差距越大:

注意力运行时间与显存占用(论文图 3):左图为前向+反向运行时间,右图为显存占用——标准注意力显存随序列长度平方增长,FlashAttention 线性增长(来源:arXiv:2205.14135 图 3)
注意力运行时间与显存占用(论文图 3):左图为前向+反向运行时间,右图为显存占用——标准注意力显存随序列长度平方增长,FlashAttention 线性增长(来源:arXiv:2205.14135 图 3)

加速比数据。 在 A100 上(batch 8、头维度 64、12 头),FlashAttention 相对 PyTorch 标准注意力在常见序列长度(128–2K)下普遍取得 2–4 倍加速;GPT-2 上的纯注意力前向达到 7.6 倍。论文附录给出了更多硬件上的数据(RTX 3090、T4 同样有 2–4 倍),说明收益不是 A100 专属,而是带宽-容量比例普遍失衡的必然结果。下方这张 A100 加速比图展示了不同序列长度下的具体表现:

A100 上 FlashAttention 相对 PyTorch 标准注意力的加速比(论文图 5,batch 8、头维度 64、12 头):常见序列长度下普遍 2–4 倍,长序列加速比更高(来源:arXiv:2205.14135 图 5)
A100 上 FlashAttention 相对 PyTorch 标准注意力的加速比(论文图 5,batch 8、头维度 64、12 头):常见序列长度下普遍 2–4 倍,长序列加速比更高(来源:arXiv:2205.14135 图 5)

端到端训练收益。 用 FlashAttention 训练模型,在相同显存预算下能跑更长的序列,或者跑同样配置更快:BERT-large(序列长度 512)比 MLPerf 1.1 的训练速度纪录快 15%;GPT-2 small/medium 比 HuggingFace 和 Megatron-LM 的实现快最多 3 倍;Long-Range-Arena(LRA,序列长度 1K–4K)快 2.4 倍。

长序列带来模型质量提升。 这可能是 FlashAttention 最深远的影响:显存从 O(N2)O(N^2) 降到 O(N)O(N),意味着此前跑不动的长序列现在能跑了。GPT-2 在 4 倍上下文长度下困惑度(perplexity)改善 0.7;长文档分类任务上,建模更长序列带来 6.4 个点的 F1 提升;更标志性的是 Path-X(序列长度 16K)——这是第一个在该任务上超过随机水平的 Transformer 模型。今天大模型动辄 128K 的上下文,起点就在这里。

局限:为什么 FlashAttention-1 只达到 25–40% 的峰值利用率#

FlashAttention-1 解决了”内存瓶颈”,但它自己也不完美:在 A100 上,它只达到理论峰值 FLOPs 的 25–40%(FlashAttention-2 论文中的评估)。换句话说,它不再受限于 HBM,却受限于自己对 GPU 并行模型的使用方式。三个原因:

第一,并行粒度受限。 FlashAttention-1 把并行只放在 batch 维和 head 维:每个注意力头用一个线程块(thread block)处理,总共 batch × head 个线程块。A100 有 108 个 SM,当 batch 和 head 数的乘积小于 108(长序列场景下 batch 通常很小,比如 1–8),大量 SM 空闲,占用率(occupancy)不足。序列长度维本来是 FlashAttention 里最大的可并行维度,却没有被利用。

第二,warp 之间的分工不合理。 线程块内通常用 4 个 warp 协作:FlashAttention-1 把 KKVV 分给不同 warp(split-K 方案),每个 warp 算出 QKTQK^T 的一个切片后,需要把结果写进共享内存、同步、再归约相加,才能继续 PVPV 的乘法——这引入了不必要的共享内存读写和同步开销。

第三,非矩阵乘操作(non-matmul FLOPs)占比偏高。 softmax 的重缩放、归一化等逐元素操作挤占了本可以留给矩阵乘的指令周期。

这三个问题正好是 FlashAttention-2 的三个改进方向:把外层循环从 KKVV 块改成 QQ 块,从而把序列维变成可并行维;把 warp 分工从”切 KKVV“改成”切 QQ“(每个 warp 独立算完自己那行块的 softmax 和输出,warp 之间零通信);以及算法层面的微调,把非矩阵乘操作降到最低。FlashAttention-2 在 A100 上把这三点修完后,达到了 50–73% 的峰值利用率,相对 FlashAttention-1 再快约 2 倍。

下一篇将讨论 FlashAttention-2 的算法微调、序列维并行与 warp 工作划分(work partitioning)的完整细节,以及针对解码(decode)阶段的 FlashDecoding 优化。

小结#

FlashAttention-1 的价值可以浓缩为三句话:

  • 目标变了:从”减少 FLOPs”变成”减少 HBM 访问”,IO 复杂度 Θ(N2d2/M)\Theta(N^2d^2/M) 替代 N2N^2 成为设计目标,并证明了这是精确注意力的下界;
  • 手段是数据流:分块 + 在线 softmax + 反向重计算,让 N×NN \times N 中间矩阵从不落盘,显存 O(N2)O(N)O(N^2) \to O(N),速度 2–7.6 倍提升,且数学上精确无近似;
  • 它开启了整个时代:长上下文训练成为可能(Path-X 从不可行到可行),而它留下的并行化短板直接催生了 FlashAttention-2/3/4 三代演进——本系列的第二篇就讲 FlashAttention-2。

今天再看 FlashAttention,它的思想已经渗透到几乎所有推理框架的底层:vLLM 的 PagedAttention 处理的是 KV 缓存的分页,FlashAttention 处理的是计算过程的分块,两者互补构成了现代长上下文推理的两个支柱。理解 FlashAttention 的 IO 感知视角,是读懂后续一切注意力优化(FlashDecoding、Triton 手写 kernel、MLA 的低秩投影等)的前提。

参考资料#

  1. FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness — FlashAttention-1 原始论文(NeurIPS 2022),本文主要依据
  2. FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning — FlashAttention-2 论文,下一篇的主要依据
  3. Online normalizer calculation for softmax — Milakov & Gimelshein 2018,在线 softmax 的原始论文
  4. FlashAttention: Making Attention I/O-Aware — Hugging Face 对 FlashAttention 的解读
  5. FlashAttention — IO Analysis and Evolution — Hugging Face 博客,含 Roofline 模型下的算术强度分析
  6. 从 Online Softmax 到 FlashAttention — 腾讯云开发者社区,在线 softmax 到 FlashAttention 的完整推导
  7. 从零开始用自定义 Triton 内核编写 FlashAttention-2 — 阿里云开发者社区,Triton 实现 FlashAttention-2 的完整教程
  8. Dao-AILab/flash-attention GitHub 仓库 — FlashAttention 官方开源实现
  9. FlashAttention 完全拆解(四):FlashAttention-4——面向 Blackwell 的算法-流水线协同设计 — 本站对 FlashAttention 系列最新一代的拆解
  10. 浮点数与数值稳定性:从 IEEE 754 到 FP8 的 LLM 精度世界 — 本站文章,含在线 softmax 的完整数值分析

文章分享

如果这篇文章对你有帮助,欢迎分享给更多人!

FlashAttention 完全拆解(一):IO 感知与分块注意力算法
https://pinghaoyang.com.cn/aigc/posts/flashattention-part-1/
作者
平昊阳
发布于
2026-08-26
许可协议
CC BY-NC-SA 4.0

评论区

Profile Image of the Author
平昊阳
乘长风,破巨浪, 展鸿图于未央!
--
总访问量
--
访客数
公告
欢迎来到我的个人博客!欢迎关注交流吖!
更多相关公告,见
社交-留言」。
音乐
封面

音乐

暂未播放

0:000:00
暂无歌词
站点统计
文章
84
分类
18
标签
108
总字数
701,808
运行时长
0
最后活动
0 天前

文章目录