音乐
暂未播放
FlashAttention 完全拆解(一):IO 感知与分块注意力算法
背景:标准注意力实现的三个瓶颈#
FlashAttention 于 2022 年由斯坦福大学的 Tri Dao 等人提出,发表在 NeurIPS 2022 上。它是过去几年对 Transformer 训练和推理影响最大的系统优化之一:这篇论文之前,注意力(attention)实现的速度和显存占用长期受制于一个简单的工程事实——中间矩阵太大,必须反复读写慢速的显存。本站的 FlashAttention 完全拆解(四) 已经讨论过这条技术路线的最终形态,本文回到起点,把 FlashAttention-1/2 这代经典设计彻底讲透。
注意力层的计算是:
Attention(Q,K,V)=softmax(dQKT)V其中 Q,K,V∈RN×d 分别是查询(query)、键(key)、值(value)矩阵,N 是序列长度,d 是每个头的维度(head dimension)。以 d 为头的维数,QKT 得到一个 N×N 的注意力分数矩阵 S,softmax 按行归一化得到概率矩阵 P,再乘上 V 得到输出 O。
在 FlashAttention 出现之前,GPU 上的标准实现(论文里称为 Algorithm 0)把这三步拆成三个独立的 kernel,每个 kernel 都从高带宽内存(High Bandwidth Memory,HBM)读入、算完、写回:
- 加载 Q、K,计算 S=QKT,把 N×N 的 S 写回 HBM;
- 重新读入 S,计算 P=softmax(S)(含行内最大值归约、指数、求和、归一化),把 P 写回 HBM;
- 重新读入 P 和 V,计算 O=PV,写回 HBM。
这个实现有三个互相纠缠的瓶颈。
瓶颈一:N×N 中间矩阵的物化(materialization)带来 O(N2) 显存占用。 序列长度每翻一倍,S 和 P 的体积翻四倍。序列长度 4096、头维度 128 时,每个矩阵在 FP16 下约 33.5 MB,S、P 合起来 67 MB;到 16384 时每个约 537 MB;到 65536 时每个约 8.6 GB,两个矩阵超过 17 GB——一张 A100 的显存就被注意力中间结果占掉一大半。这直接限制了模型能处理的上下文长度,也解释了为什么 FlashAttention 之前的主流模型都把序列长度钉死在 2K 附近。
瓶颈二:HBM 访问次数是 Θ(Nd+N2)。 标准的三个 kernel 把 S 和 P 各写一次、各读一次,加上 Q、K、V、O 的读写,总共有 Θ(N2) 量级的数据在 HBM 上进进出出。注意这里说的是访问的元素个数,单位量级是 Θ(N2) 而不是 Θ(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=4096、d=128、FP16 精度下,标准注意力完成一次前向约需要 4N2d≈8.6 GFLOP(两次矩阵乘各 2N2d,softmax 约 3N2),而 HBM 流量约 134 MB(S、P 各写读一遍占了大头),算术强度约 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×N 的 S 和 P 依然要物化、要写读,Θ(N2) 的 HBM 访问依然存在。FlashAttention 的贡献不是把某个 kernel 调快,而是从根本上改变数据流,让 Θ(N2) 这个量级消失。
核心思想:IO 感知(IO-aware)#
FlashAttention 的核心思想用一句话概括:设计注意力算法时,把 HBM 访问次数当作首要优化目标,而不是 FLOPs。
论文用”IO 感知”(IO-aware)来命名这个原则,背后是一条硬件趋势:计算速度的提升长期快于内存速度的提升,因此大量操作从计算受限滑向内存受限。对注意力而言,标准实现的 HBM 访问是 Θ(Nd+N2),FlashAttention 把它压到 Θ(N2d2/M),其中 M 是片上 SRAM 的大小——典型配置下这比 N2 少了一个数量级以上的常数因子。更关键的是,论文证明了这是精确注意力能到达的下界(lower bound,详见后文)。
实现 IO 感知需要三件套,缺一不可:
- 分块(tiling):把 Q、K、V 切成小块,让每个块的处理在片上 SRAM 内完成,N×N 的中间矩阵从头到尾不落 HBM;
- 在线 softmax(online softmax):softmax 需要全局的统计量(行最大值与归一化和),分块后只能看到局部,必须用流式算法在遍历过程中维护并修正这些统计量;
- 重计算(recomputation):反向传播需要 S、P 这两个 N×N 的中间矩阵,FlashAttention 不存储它们,而是在反向时从块和统计量中重新算出来。
要强调的一点是:FlashAttention 是精确算法(exact algorithm),不是近似——它和标准注意力在数学上完全等价,唯一的差异是浮点运算顺序不同带来的舍入误差。这与当时流行的稀疏注意力(sparse attention)、线性注意力(linear attention)等近似方法有本质区别,也是它能成为事实标准的最重要原因:不需要任何精度权衡,就能白拿加速和显存节省。
机制一:分块(Tiling)与块大小设计#
分块的目标是让”一个 Q 块 + 一个 K 块 + 一个 V 块 + 对应的分数块 Sij“能同时放进片上 SRAM。设 SRAM 大小为 M(A100 上约 192 KB),则块大小取:
Bc=⌈4dM⌉,Br=min(⌈4dM⌉,d)其中 Bc 是 K、V 块的”行数”(每个块是 Bc×d),Br 是 Q 块的”行数”(每个块是 Br×d)。为什么是除以 4?因为 SRAM 里要同时驻留四样东西:Qi(Br×d)、Kj(Bc×d)、Vj(Bc×d)、以及中间结果 Sij=QiKjT(Br×Bc),四份都占 Θ(M) 的量级,所以每份只能拿到 M/4。三个约束条件分别是:
Bc⋅d=O(M),Br⋅d=O(M),Br⋅Bc=O(M)前两个保证 Kj、Vj 和 Qi 装得下,第三个保证它们相乘出来的分数块 Sij 也装得下。联立解出来就是上面的取值:Bc 取 Θ(M/d),Br 取 min(Θ(M/d),d)——注意 Br 还受 d 本身限制,因为 Br×Bc 的约束要求 Br=O(M/Bc)=O(d)。
代入 A100 的实际数字看看(以 FP16 元素计:192 KB ÷ 2 字节/元素 = 98304 个元素,d=128):M/(4d)=98304/512=192,所以 Bc=192,Br=min(192,128)=128。也就是说,Q 块取 128×128,K、V 块取 192×128,分数块是 128×192——四个块合计 128×128+192×128+192×128+128×192=90112 个元素,约 180 KB,刚好塞进 192 KB 的 SRAM,还留出一点余量给掩码和 dropout 状态。
这里要说明一个工程细节:上面的公式是论文给出的理论取值,实际 kernel 里块大小还要扣除掩码、dropout 状态、寄存器等开销,而论文的约束 O(M) 是渐近意义下的,常数因子由实现取舍。因此实际 FlashAttention kernel 常用 Br=Bc=64 或 128 这样的取值,论文在块大小选择上的态度是”在寄存器溢出和 SRAM 容量之间手工调优”。理论公式的意义在于揭示块大小只依赖 M 和 d,不依赖序列长度 N——这保证了不管序列多长,每个块的片上占用都是常数,IO 复杂度才能只由 M 决定。这个推导逻辑比具体数字更重要。
分块之后,Q 被分成 Tr=⌈N/Br⌉ 个块,K、V 各被分成 Tc=⌈N/Bc⌉ 个块。FlashAttention-1 的循环组织是:外层循环遍历 K、V 块,内层循环遍历 Q 块——每加载一个 Kj、Vj 块进 SRAM,就把它和所有 Q 块各算一遍(论文图 1 的红色箭头是外层,蓝色箭头是内层):

图中虚线框代表那个 N×N 的注意力矩阵——在 FlashAttention 里它从头到尾不存在于任何地方。Sij 算出来后在 SRAM 里直接参与 softmax 和加权求和,用完后被下一个 Si,j+1 覆盖。同一张图的右边是作者在 GPT-2 上测的加速数据:仅注意力部分就比 PyTorch 标准实现快 7.6 倍。
机制二:在线 softmax——分块的数学支柱#
分块带来一个直接的数学障碍:softmax 是全局操作,而分块之后每个 Sij 只包含部分分数。
回忆 softmax 的定义:P 的第 i 行是所有 eSi,j 除以该行的和 ∑jeSi,j。这个”行和”需要看完整行才能算出来;为了数值稳定,还要先减去行最大值 mi=maxjSi,j,否则指数会溢出。标准实现要两趟(先求最大值、再求指数和),而分块流式计算要求一趟遍历完成所有工作——处理某个 K、V 块时,后面的块还没看到,无法知道这一行的最大值和总和。
解决办法是维护两个随遍历推进的累计量:运行最大值(running max)mi 和运行归一化常数 ℓi。处理第 j 个块时:
minew=max(mi,m~ij),m~ij=rowmax(Sij)ℓinew=emi−minewℓi+em~ij−minewℓ~ij,ℓ~ij=rowsum(eSij−m~ij)其中 mi、ℓi 是前 j−1 个块累积下来的统计量(以”当时的最大值”为基准),m~ij、ℓ~ij 是当前块 Sij 自己的行最大值和”以本块最大值为基准”的指数和。注意 ℓi 的定义细节:它是以当前最大值 mi 为基准的指数和——这是整个算法能保持精确的关键。当新块的最大值 m~ij 超过旧值 mi 时,之前累积的所有项都要乘上 emi−minew 这个缩放因子(rescaling factor),把基准从旧最大值切换到新最大值;而当最大值没变时,缩放因子是 e0=1,退化为普通累加。
输出累加器 Oi 的更新同样带重缩放:
Oi←diag(ℓinew)−1(diag(ℓi)emi−minewOi+em~ij−minewP~ijVj)其中 P~ij=exp(Sij−m~ij) 是本块的”未归一化概率”块。所有块处理完后,做最后一次归一化 Oi←Oi/ℓi,得到的就是精确的 softmax(QiKT)V。
为什么这套递归在数学上精确?可以用归纳不变式(inductive invariant)理解:每一步结束时,不变式”Oi/ℓi = 前 j 个块对应的精确 softmax 加权和”始终成立。当最大值不更新时只是普通累加;当最大值增大时,旧项和新项被同一个标量重缩放——同一行的所有项始终除以同一个 ℓi,比例关系不变,所以最终结果与”先看完全部行再归一化”完全一致。算法做的只是把除法推迟到了最后一步。
这套机制在本站的《浮点数与数值稳定性》一文中有完整的推导和数值分析(包括为什么它不会溢出、为什么下溢只发生在”本来就该被忽略”的项上),这里不再重复。值得单独强调的直觉是:在线 softmax 让所有中间值始终保持在 [0,1] 量级——这正是它数值稳定的原因,也是分块注意力能存在的数学前提。后面 FlashAttention-4 的”条件重缩放”、DeepSeek 的 FP8 注意力等工程优化,都是在这套机制上做加减法。
完整前向算法#
把三件套合起来,就是论文的 Algorithm 1(FlashAttention 前向传播):
1输入:Q, K, V ∈ R^{N×d} 位于 HBM,片上 SRAM 大小为 M2 1: 块大小 B_c = ⌈M/(4d)⌉,B_r = min(⌈M/(4d)⌉, d)3 2: 在 HBM 初始化 O = (0)_{N×d},ℓ = (0)_N,m = (-∞)_N4 3: 把 Q 分成 T_r = ⌈N/B_r⌉ 个块;把 K、V 各分成 T_c = ⌈N/B_c⌉ 个块5 4: 把 O、ℓ、m 分成对应的 T_r 个块6 5: for j = 1 to T_c do7 6: 加载 K_j、V_j 到 SRAM8 7: for i = 1 to T_r do9 8: 加载 Q_i、O_i、ℓ_i、m_i 到 SRAM10 9: 在片上计算 S_ij = Q_i K_j^T ∈ R^{B_r×B_c}1110: 在片上计算 m̃_ij = rowmax(S_ij),P̃_ij = exp(S_ij − m̃_ij),ℓ̃_ij = rowsum(P̃_ij)1211: 在片上计算 m_i^new = max(m_i, m̃_ij),ℓ_i^new = e^{m_i−m_i^new}·ℓ_i + e^{m̃_ij−m_i^new}·ℓ̃_ij1312: 更新 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),写回 HBM1413: 把 ℓ_i、m_i 更新为 ℓ_i^new、m_i^new,写回 HBM1514: end for1615: end for1716: 返回 O逐行看这个算法的关键决策:
- 第 1 行的块大小只依赖 M 和 d,前面已经推导过;
- 第 9 行是两处矩阵乘中的第一处:Sij=QiKjT 在张量核心上完成,这是注意力计算中效率最高的部分;
- 第 10–11 行是在线 softmax 的状态更新,全部在 SRAM 内完成,Sij 算完即被消费;
- 第 12 行是第二次矩阵乘:P~ijVj,把概率块和值块相乘累加进输出。注意这里 Oi 在每次内层迭代都写回 HBM——这是论文表述层面的简化,实际 kernel 里 Oi 留在 SRAM/寄存器中,整行 i 处理完才写回一次;
- 第 13 行:mi、ℓi 这两个 Br 维的小向量是唯一需要跨块传递的状态——它们记录了”到目前为止看了多少”,是重计算(反向传播)时需要保存的全部统计信息。
整体复杂度(论文定理 1):算法使用 O(N2d) 次浮点运算(和标准实现同量级),除输入输出外只需要 O(N) 的额外内存——m 和 ℓ 各 N 个标量。显存从 O(N2) 降到 O(N),这就是”线性注意力内存”的由来,且不牺牲任何精度。
IO 复杂度分析:从 Θ(Nd+N2) 到 Θ(N2d2/M)#
论文把 HBM 访问量的分析形式化为 IO 复杂度(IO complexity),这是理解 FlashAttention 为什么快、以及为什么”已经到极限”的关键。
标准注意力的 HBM 访问是 Θ(Nd+N2)。 逐 kernel 数:S=QKT 一步读 Q、K(Nd 量级)写 S(N2 量级);softmax 一步读 S 写 P(各 N2);O=PV 一步读 P、V 写 O(N2+Nd)。三项相加,主导项是 Θ(N2)——序列长度平方增长,和显存占用一样。
FlashAttention 的 HBM 访问是 Θ(N2d2/M)。 逐块数:每个 K、V 元素从 HBM 只被加载一次(第 6 行,外层循环遍历时加载,之后全部复用);而 Q 和 O 每个外层块都要被完整读一遍,即总共被读 Tc=N/Bc 次,每趟读 O(Nd) 个元素,合计 O(Nd⋅Tc)=O(N2d/Bc)=O(N2d2/M)(代入 Bc=Θ(M/d))。K、V 各一次读入是 Θ(Nd),被 N2d2/M 主导。
两个量级对比一下:标准是 N2,FlashAttention 是 N2d2/M。在典型配置(d=64–128,M≈100 KB 量级)下,d2≪M,所以 FlashAttention 的 HBM 访问比标准实现少一个很大的常数因子——论文在 A100 上实测最多减少约 9 倍。这个因子不随 N 增长而改变,所以长序列下收益是持续的:无论是 2K 还是 64K 上下文,HBM 流量都从”随 N2 爆炸”变成”随 N2 增长但系数小得多”。
更重要的是下界:精确注意力的 IO 复杂度已经渐近最优。 论文的命题 3(Proposition 3)证明:对 M∈[d,Nd] 范围内的任意 SRAM 大小,不存在任何精确注意力算法能以 o(N2d2/M) 的 HBM 访问量完成计算。证明思路是取极端情形 M=Θ(Nd)(SRAM 大到能装下整个 K、V),此时 N2d2/M=Θ(Nd),而任何算法至少要把 Q、K、V、O 读入写出一次——这是物理下限,任何算法都逃不掉。
这个下界有深远的实践含义:FlashAttention-2/3/4 的改进空间不在 IO 复杂度本身(那已经是下界),而在常数因子、并行度、硬件利用率。后文会看到,FlashAttention-2 的加速正是这样来的——算法没变,变的是线程块和 warp 怎么分工。
反向传播:用重计算(recomputation)换显存#
前向的分块解决了 S、P 的物化问题,但反向传播(backward pass)还藏着一个 O(N2) 的坑:计算梯度需要 S 和 P。具体来说,设 dO 是输出的梯度,标准反演需要:
dV=PTdO,dP=dO⋅VT,dS=P∘(dP−D),dQ=dS⋅K,dK=dST⋅Q其中 Di=rowsum(dOi∘Oi) 是 softmax 反向的修正项,∘ 表示逐元素乘。如果像标准实现那样把 S、P 存下来给反向用,显存又回到 O(N2),前向省下的空间全部白费。
FlashAttention 的选择是不存 S、P,反向时重算。前向结束时,每个 Q 块只需把两样东西留给反向:输出块 Oi,以及统计量 mi、ℓi(合计 O(N) 个标量)。反向的每个内层迭代从 HBM 加载 Qi、Kj、Vj、Oi、dOi 到 SRAM,在片上重算 Sij 和概率块:
Pij=diag(ℓi)−1exp(Sij−mi)然后用上面的五个公式在块内完成 dVj、dKj、dQi 的梯度累积。重算的成本只是多一次 QiKjT 的矩阵乘——而这本来就是张量核心最擅长的操作,且在 SRAM 内完成,不产生额外 HBM 流量。
论文把这种设计归类为选择性梯度检查点(selective gradient checkpointing):标准的梯度检查点技术用重算换内存,代价是训练变慢;FlashAttention 的反向虽然 FLOPs 更多(约翻倍),但 HBM 访问从标准反向的 Θ(Nd+N2) 降到和正向一样的 Θ(N2d2/M)——内存少了、FLOPs 多了、反而更快,因为反向同样是内存受限的。论文定理 5 给出了与正向对称的复杂度结论。这大概是”以算换访”在系统优化史上最成功的一次应用。
块稀疏扩展(Block-Sparse FlashAttention)#
论文还顺带把 FlashAttention 扩展成了块稀疏版本(block-sparse FlashAttention):给定一个块级掩码 M~∈{0,1}N×N,跳过被掩码盖住的 Sij 块的计算,IO 复杂度按稀疏度比例降低——比如稀疏度为 1/s 时 HBM 访问约减少 s 倍。
这个扩展的学术意义在于:它证明了 FlashAttention 可以作为一种通用原语,让精确注意力和各种近似注意力(稀疏、低秩、核方法)在同一套 IO 框架下比较和实现。本站在《BLASST》里拆解过的动态块稀疏注意力,本质上就是这条线的延续——把”掩码从哪里来”从静态变成运行时决策。不过 FlashAttention 的块稀疏版本由于掩码固定的限制,在真实场景中逐渐被更灵活的方案取代,这里就不展开细讲了。
性能表现#
论文在多个 GPU 和多个模型上做了评测,几组关键数据:
HBM 访问是运行时间的决定因素。 论文图 2 用 GPT-2 medium(序列长度 1024、头维度 64、16 头、batch 64)在 A100 上做微基准:左图对比标准注意力和 FlashAttention 的前向+反向运行时间,中图展示 FlashAttention 前向时间随 HBM 访问减少而下降——访问量越少越快,直到某个拐点(约 9 倍访问量差异):

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

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

端到端训练收益。 用 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),意味着此前跑不动的长序列现在能跑了。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 把 K、V 分给不同 warp(split-K 方案),每个 warp 算出 QKT 的一个切片后,需要把结果写进共享内存、同步、再归约相加,才能继续 PV 的乘法——这引入了不必要的共享内存读写和同步开销。
第三,非矩阵乘操作(non-matmul FLOPs)占比偏高。 softmax 的重缩放、归一化等逐元素操作挤占了本可以留给矩阵乘的指令周期。
这三个问题正好是 FlashAttention-2 的三个改进方向:把外层循环从 K、V 块改成 Q 块,从而把序列维变成可并行维;把 warp 分工从”切 K、V“改成”切 Q“(每个 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) 替代 N2 成为设计目标,并证明了这是精确注意力的下界;
- 手段是数据流:分块 + 在线 softmax + 反向重计算,让 N×N 中间矩阵从不落盘,显存 O(N2)→O(N),速度 2–7.6 倍提升,且数学上精确无近似;
- 它开启了整个时代:长上下文训练成为可能(Path-X 从不可行到可行),而它留下的并行化短板直接催生了 FlashAttention-2/3/4 三代演进——本系列的第二篇就讲 FlashAttention-2。
今天再看 FlashAttention,它的思想已经渗透到几乎所有推理框架的底层:vLLM 的 PagedAttention 处理的是 KV 缓存的分页,FlashAttention 处理的是计算过程的分块,两者互补构成了现代长上下文推理的两个支柱。理解 FlashAttention 的 IO 感知视角,是读懂后续一切注意力优化(FlashDecoding、Triton 手写 kernel、MLA 的低秩投影等)的前提。
参考资料#
- FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness — FlashAttention-1 原始论文(NeurIPS 2022),本文主要依据
- FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning — FlashAttention-2 论文,下一篇的主要依据
- Online normalizer calculation for softmax — Milakov & Gimelshein 2018,在线 softmax 的原始论文
- FlashAttention: Making Attention I/O-Aware — Hugging Face 对 FlashAttention 的解读
- FlashAttention — IO Analysis and Evolution — Hugging Face 博客,含 Roofline 模型下的算术强度分析
- 从 Online Softmax 到 FlashAttention — 腾讯云开发者社区,在线 softmax 到 FlashAttention 的完整推导
- 从零开始用自定义 Triton 内核编写 FlashAttention-2 — 阿里云开发者社区,Triton 实现 FlashAttention-2 的完整教程
- Dao-AILab/flash-attention GitHub 仓库 — FlashAttention 官方开源实现
- FlashAttention 完全拆解(四):FlashAttention-4——面向 Blackwell 的算法-流水线协同设计 — 本站对 FlashAttention 系列最新一代的拆解
- 浮点数与数值稳定性:从 IEEE 754 到 FP8 的 LLM 精度世界 — 本站文章,含在线 softmax 的完整数值分析
文章分享
如果这篇文章对你有帮助,欢迎分享给更多人!
部分内容可能已过时
评论区
分享你的想法,与大家交流讨论
音乐
暂未播放



