FlashAttention 1~4 全解汇总

2183 字
11 分钟
FlashAttention 1~4 全解汇总

Datawhale
Datawhale

大家好,我是芯缘,是 Datawhale 社区发起的 2026 年 8 月“llm-algo-leetcode 推理优化方向”组队学习活动的运营助教。本文记录了我学习 Task 2: Attention 访存瓶颈FlashAttention 的笔记。

FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness(2022)#

一种具有 IO 感知能力的快速、显存高效的精确注意力算法。

Algorithm#

Background#

(1) 内存的金字塔结构,HBM 带宽远小于 On-chip SRAM。

(2) SSPP 均为 O(N2)O(N^2) ,被反复读写 HBM 时,即使 Attention 的计算量为 O(N2d)O(N^2d),整体容易呈现 Memory-Bound

Solution#

沿序列长度分块,在 SRAM 中完成局部 score、online softmax 和输出累加,从而避免将完整的 SSPP 写入 HBM。

Input: Q,K,VRN×d,SRAM size=MOutput: O=softmax(QK)V// Step 1: Set block sizeBcM4d,Brmin ⁣(M4d,d)// Step 2: InitializeO0,0,m// Step 3: Split by sequence lengthQ{Q1,,QTr},QiRBr×dK{K1,,KTc},KjRBc×dV{V1,,VTc},VjRBc×dO{O1,,OTr},{1,,Tr},m{m1,,mTr}// Step 4: Forward loopfor j1 to Tc doKj,Vj: HBM  SRAMfor i1 to Tr doQi,Oi,i,mi: HBM  SRAMSijQiKjminewmax ⁣(mi,rowmax(Sij))P~ijexp ⁣(Sijminew)inewemiminewi+rowsum ⁣(P~ij)OinewemiminewiOi+P~ijVjinewOi,i,miOinew,inew,minewOi,i,mi: SRAM  HBMend forend for// Step 5: Returnreturn O\begin{aligned} \textbf{Input: }& Q,K,V\in\mathbb{R}^{N\times d},\quad \text{SRAM size}=M \\[4pt] \textbf{Output: }& O=\operatorname{softmax}(QK^\top)V \\[8pt] &\text{// Step 1: Set block size} \\[2pt] &B_c \leftarrow \left\lceil \frac{M}{4d} \right\rceil,\quad B_r \leftarrow \min\!\left(\left\lceil \frac{M}{4d} \right\rceil,d\right) \\[8pt] &\text{// Step 2: Initialize} \\[2pt] &O \leftarrow 0,\quad \ell \leftarrow 0,\quad m \leftarrow -\infty \\[8pt] &\text{// Step 3: Split by sequence length} \\[2pt] &Q \rightarrow \{Q_1,\dots,Q_{T_r}\},\quad Q_i\in\mathbb{R}^{B_r\times d} \\[4pt] &K \rightarrow \{K_1,\dots,K_{T_c}\},\quad K_j\in\mathbb{R}^{B_c\times d} \\[4pt] &V \rightarrow \{V_1,\dots,V_{T_c}\},\quad V_j\in\mathbb{R}^{B_c\times d} \\[4pt] &O \rightarrow \{O_1,\dots,O_{T_r}\},\quad \ell \rightarrow \{\ell_1,\dots,\ell_{T_r}\},\quad m \rightarrow \{m_1,\dots,m_{T_r}\} \\[8pt] &\text{// Step 4: Forward loop} \\[2pt] &\textbf{for } j \leftarrow 1 \textbf{ to } T_c \textbf{ do} \\[2pt] &\quad K_j,V_j:\text{ HBM }\rightarrow\text{ SRAM} \\[6pt] &\quad \textbf{for } i \leftarrow 1 \textbf{ to } T_r \textbf{ do} \\[2pt] &\quad\quad Q_i,O_i,\ell_i,m_i:\text{ HBM }\rightarrow\text{ SRAM} \\[6pt] &\quad\quad S_{ij}\leftarrow Q_iK_j^\top \\[6pt] &\quad\quad m_i^{\text{new}} \leftarrow \max\!\left(m_i,\operatorname{rowmax}(S_{ij})\right) \\[6pt] &\quad\quad \widetilde{P}_{ij} \leftarrow \exp\!\left(S_{ij}-m_i^{\text{new}}\right) \\[6pt] &\quad\quad \ell_i^{\text{new}} \leftarrow e^{m_i-m_i^{\text{new}}}\ell_i +\operatorname{rowsum}\!\left(\widetilde{P}_{ij}\right) \\[6pt] &\quad\quad O_i^{\text{new}} \leftarrow \frac{ e^{m_i-m_i^{\text{new}}}\ell_iO_i+\widetilde{P}_{ij}V_j }{ \ell_i^{\text{new}} } \\[10pt] &\quad\quad O_i,\ell_i,m_i \leftarrow O_i^{\text{new}},\ell_i^{\text{new}},m_i^{\text{new}} \\[4pt] &\quad\quad O_i,\ell_i,m_i:\text{ SRAM }\rightarrow\text{ HBM} \\[2pt] &\quad \textbf{end for} \\[2pt] &\textbf{end for} \\[6pt] &\text{// Step 5: Return} \\[2pt] &\textbf{return } O \end{aligned}

Details#

Block Size设置 关键约束是一次计算时,Kj,VjK_j,V_jQi,OiQ_i,O_i 和局部 score block SijS_{ij} 都要放进 On-chip SRAM:

Bcd=O(M),Brd=O(M),BrBc=O(M).B_cd=O(M),\qquad B_rd=O(M),\qquad B_rB_c=O(M).

其中 Bcd=O(M)B_cd=O(M) 来自 Kj,VjK_j,V_j 的驻留需求,Brd=O(M)B_rd=O(M) 来自 Qi,OiQ_i,O_i 的驻留需求,BrBc=O(M)B_rB_c=O(M) 来自 Sij=QiKjRBr×BcS_{ij}=Q_iK_j^\top\in\mathbb{R}^{B_r\times B_c} 的临时存储需求。

FA1 让 BcB_c 尽量大:

Bc=Θ(Md),B_c=\Theta\left(\frac{M}{d}\right),

这样外层 K,VK,V block 数更少,Q,OQ,O 被重复扫描的轮数也更少。但 BcB_c 变大后,SijS_{ij} 的约束会限制 BrB_r

BrBc=O(M)Br=O(d).B_rB_c=O(M) \quad\Rightarrow\quad B_r=O(d).

因此论文中取:

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).

为什么 mm 初始化为 -\infty mm 记录的是当前已经扫描过的 score 的行最大值。online softmax 在处理第一个 SijS_{ij} block 时,希望新最大值完全由这个 block 决定:

minew=max ⁣(mi,rowmax(Sij))=rowmax(Sij).m_i^{\text{new}} =\max\!\left(m_i,\operatorname{rowmax}(S_{ij})\right) =\operatorname{rowmax}(S_{ij}).

mim_i 初始化成 -\infty,就相当于告诉算法“当前还没有看过任何 score”。同时旧分母和旧输出的缩放项会自然消失:

emiminewi=eminew0=0.e^{m_i-m_i^{\text{new}}}\ell_i =e^{-\infty-m_i^{\text{new}}}\cdot 0 =0.

这样第一块和后续块可以使用完全相同的更新公式,不需要单独写 first block 的特殊分支。

复杂度分析 计算量仍然由两次矩阵乘主导,总体 FLOPs 仍为:

O(N2d).O(N^2d).

HBM 访问量的变化更关键。所有 K,VK,V block 在外层循环中合计读入一次,代价是 Θ(Nd)\Theta(Nd);每处理一个 K,VK,V block,需要从 HBM 扫描一遍全部 Q,OQ,O block 以及 m,m,\ell 统计量,主导代价是 Θ(Nd)\Theta(Nd),不是 Θ(Nd+N2)\Theta(Nd+N^2)

这里不会出现每轮 Θ(N2)\Theta(N^2) 的 HBM 访问,因为局部 SijS_{ij}P~ij\widetilde{P}_{ij} 虽然总体计算量覆盖了 N2N^2 个 attention score,但它们只在 SRAM 中短暂生成和消费,并不会作为完整矩阵写入 HBM。因为:

Tc=NBc=Θ(NdM),T_c=\frac{N}{B_c} =\Theta\left(\frac{Nd}{M}\right),

所以 HBM 访问量为:

Θ(Nd+NdTc)=Θ(NdTc)=Θ(N2d2M).\Theta(Nd+NdT_c) =\Theta(NdT_c) =\Theta\left(\frac{N^2d^2}{M}\right).

标准 Attention 需要把完整 SSPP 写入并读出 HBM,HBM 访问量为:

Θ(Nd+N2).\Theta(Nd+N^2).

当片上 SRAM 足够容纳较大的 block,且通常 Md2M\gg d^2 时,FlashAttention-1 可以显著减少 HBM 访问。 注:SRAM容量通常为约 6464164 KiB164\ \text{KiB};FP16/BF16 下约 32K32\text{K}82K82\text{K} 个元素。

Source Code#

【这一部分内容待完善~】

FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning(2023)#

通过更好的并行性与工作划分实现更快的注意力算法。

Algorithm#

Background#

(1) 矩阵乘由 Tensor Core 加速,而 non-matmul FLOPs 更贵。FA1 的Thread Block 内部的 Non-Matmul 计算偏多。

(2) FA1 的并行划分局限于batch和head,且在短 Batch / 长序列场景下 Occupancy 不高。

Solution#

FA2 保留 FA1 的精确分块注意力思想,

  • 通过内外循环交换把并行维度扩展到序列长度方向,
  • 通过online softmax 更新公式改写实现减少每个 block 内的 non-matmul 操作。
Input: Q,K,VRN×d,Br,BcOutput: O=softmax(QK)V,L// Step 1: Split by sequence lengthQ{Q1,,QTr},QiRBr×dK{K1,,KTc},KjRBc×dV{V1,,VTc},VjRBc×dO{O1,,OTr},L{L1,,LTr}// Step 2: Parallelize over query row blocksparallel for i1 to Tr doQi: HBM  SRAMO~i(0)0,i(0)0,mi(0)// Step 3: Scan key/value blocksfor j1 to Tc doKj,Vj: HBM  SRAMSi(j)QiKjmi(j)max ⁣(mi(j1),rowmax(Si(j)))P~i(j)exp ⁣(Si(j)mi(j))i(j)emi(j1)mi(j)i(j1)+rowsum ⁣(P~i(j))O~i(j)emi(j1)mi(j)O~i(j1)+P~i(j)Vjend for// Step 4: Normalize once and save logsumexpOiO~i(Tc)i(Tc)Limi(Tc)+log ⁣(i(Tc))Oi,Li: SRAM  HBMend for// Step 5: Returnreturn O,L\begin{aligned} \textbf{Input: }& Q,K,V\in\mathbb{R}^{N\times d},\quad B_r,B_c \\[4pt] \textbf{Output: }& O=\operatorname{softmax}(QK^\top)V,\quad L \\[8pt] &\text{// Step 1: Split by sequence length} \\[2pt] &Q \rightarrow \{Q_1,\dots,Q_{T_r}\},\quad Q_i\in\mathbb{R}^{B_r\times d} \\[4pt] &K \rightarrow \{K_1,\dots,K_{T_c}\},\quad K_j\in\mathbb{R}^{B_c\times d} \\[4pt] &V \rightarrow \{V_1,\dots,V_{T_c}\},\quad V_j\in\mathbb{R}^{B_c\times d} \\[4pt] &O \rightarrow \{O_1,\dots,O_{T_r}\},\quad L \rightarrow \{L_1,\dots,L_{T_r}\} \\[8pt] &\text{// Step 2: Parallelize over query row blocks} \\[2pt] &\textbf{parallel for } i \leftarrow 1 \textbf{ to } T_r \textbf{ do} \\[2pt] &\quad Q_i:\text{ HBM }\rightarrow\text{ SRAM} \\[4pt] &\quad \widetilde{O}_i^{(0)} \leftarrow 0,\quad \ell_i^{(0)} \leftarrow 0,\quad m_i^{(0)} \leftarrow -\infty \\[8pt] &\quad \text{// Step 3: Scan key/value blocks} \\[2pt] &\quad \textbf{for } j \leftarrow 1 \textbf{ to } T_c \textbf{ do} \\[2pt] &\quad\quad K_j,V_j:\text{ HBM }\rightarrow\text{ SRAM} \\[6pt] &\quad\quad S_i^{(j)} \leftarrow Q_iK_j^\top \\[6pt] &\quad\quad m_i^{(j)} \leftarrow \max\!\left(m_i^{(j-1)},\operatorname{rowmax}(S_i^{(j)})\right) \\[6pt] &\quad\quad \widetilde{P}_i^{(j)} \leftarrow \exp\!\left(S_i^{(j)}-m_i^{(j)}\right) \\[6pt] &\quad\quad \ell_i^{(j)} \leftarrow e^{m_i^{(j-1)}-m_i^{(j)}}\ell_i^{(j-1)} +\operatorname{rowsum}\!\left(\widetilde{P}_i^{(j)}\right) \\[6pt] &\quad\quad \widetilde{O}_i^{(j)} \leftarrow e^{m_i^{(j-1)}-m_i^{(j)}}\widetilde{O}_i^{(j-1)} +\widetilde{P}_i^{(j)}V_j \\[2pt] &\quad \textbf{end for} \\[8pt] &\quad \text{// Step 4: Normalize once and save logsumexp} \\[2pt] &\quad O_i \leftarrow \frac{\widetilde{O}_i^{(T_c)}}{\ell_i^{(T_c)}} \\[6pt] &\quad L_i \leftarrow m_i^{(T_c)}+\log\!\left(\ell_i^{(T_c)}\right) \\[6pt] &\quad O_i,L_i:\text{ SRAM }\rightarrow\text{ HBM} \\[2pt] &\textbf{end for} \\[8pt] &\text{// Step 5: Return} \\[2pt] &\textbf{return } O,L \end{aligned}

Details#

通过交换内外层循环带来的收益

(1) 只需要在外循环里内循环外对 FA1 OiO_i进行归一化的rescale。

(2) 因为 QQ 在外循环,所以不需要每次都把临时的O,m,O,m,\ell进行FA1那样内外循环读写,而是每次外循环都是赋值0。

(3)在forward过程中维护的用于训练的参数m,m,\ell不需要保存到 HBM,而是只保存 logsumexp:Li=mi+log(i)L_i=m_i+\log(\ell_i) 注:如果是纯推理的话,m,,Lm,\ell,L其实都不需要写回HBM,LL推理中都不需要有,但是m,m,\ell无论训推无论FA version,都需要中途维护。

复杂度分析

总体 FLOPs 仍为:O(N2d).O(N^2d).

如果按 FA2 的 row-block outer loop 粗略估算 HBM 访问量,每个 QiQ_i block 会扫描一遍完整 K,VK,V

Tr=NBr.T_r=\frac{N}{B_r}.Tr=NBr,HBM reads of K,VΘ(NdTr).T_r=\frac{N}{B_r},\qquad \text{HBM reads of }K,V\approx \Theta(NdT_r).

HBM 访问估算:

Θ(NdTr)=Θ(N2dBr).\Theta(NdT_r) =\Theta\left(\frac{N^2d}{B_r}\right).

Source Code#

包括Work Partitioning Between Warps 【这一部分内容待完善~】

FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision(2024)#

一种利用异步执行与低精度计算实现快速且准确注意力的算法。

FlashAttention 3 的算法架构与 FlashAttention 2 一致,主要针对 Hopper (H100) 架构做了硬件级优化。

各种优化#

FA3没有算法上面的创新,因此暂缓,以后补~ Datawhale教程里面指出:

  • 核心创新 1:WGMMA 异步计算。利用 Warp Group 级指令,让 Tensor Core 在后台异步执行。
  • 核心创新 2:TMA(Tensor Memory Accelerator)。使用硬件级搬运器把数据从全局显存搬到共享内存,释放搬运线程。
  • 核心创新 3:2-Stage to Ping-Pong Pipeline。通过更高效的软件流水线掩盖访存延迟,实现计算与访存的重叠。

【这一部分内容待完善~】

FlashAttention-4: Algorithm and Kernel Pipelining Co-Design for Asymmetric Hardware Scaling(2026)#

一种面向非对称硬件扩展的算法与 Kernel 流水线协同设计方法。

Algorithm#

Background#

相比 Hopper,Blackwell 的 Tensor Core 吞吐继续提升,但 shared memory 带宽、指数函数单元(MUFU)和普通 ALU 没有同比例增长,瓶颈从单纯的矩阵乘转移到 softmax、指数计算和 shared memory traffic。

Conditional Softmax Rescaling#

FA1/FA2 的 online softmax 在每个 block 都会根据新的 running max 缩放旧输出:

Oj=emj1mjOj1+eSjmjVj.O_j=e^{m_{j-1}-m_j}O_{j-1}+e^{S_j-m_j}V_j.

FA4 观察到,只有当新的最大值明显变大时,旧输出才必须立即 rescale。因此它区分真实候选最大值 mjm_j 和当前用于缩放的基准 mˉj1\bar{m}_{j-1},只在:

mjmˉj1>τm_j-\bar{m}_{j-1}>\tau

时才更新缩放基准并执行缩放:

mˉjmj,Oj=emˉj1mˉjOj1+eSjmˉjVj.\bar{m}_j\leftarrow m_j,\qquad O_j=e^{\bar{m}_{j-1}-\bar{m}_j}O_{j-1}+e^{S_j-\bar{m}_j}V_j.

否则不更新用于缩放的 mˉ\bar{m},也不对旧输出做 rescale:

mˉjmˉj1,Oj=Oj1+eSjmˉj1Vj.\bar{m}_j\leftarrow \bar{m}_{j-1},\qquad O_j=O_{j-1}+e^{S_j-\bar{m}_{j-1}}V_j.

阈值 τ\tau 控制可以容忍的 slack。论文中典型取 τ=log2(256)=8\tau=\log_2(256)=8,对应“缩放因子不超过 256256 时先不 rescale”。如果用自然指数 exe^x 的写法理解,对应阈值就是 ln(256)\ln(256)。最后再用最终的 true normalizer 统一归一化,减少中间反复缩放带来的 non-matmul 开销。

Software-Emulated Exponential#

Blackwell 上 Tensor Core 吞吐增长很快,但指数函数单元(MUFU)没有同比例增长,softmax 里的 exp 会成为瓶颈。FA4 因此用 FMA 单元分担一部分指数计算。

softmax 数学上使用 exe^x,但实现中可以先转成 base-2 exponent:

ex=2xlog2e.e^x=2^{x\log_2 e}.

令:

y=xlog2e,yint=y,yfrac=yyint[0,1).y=x\log_2 e,\qquad y_{\text{int}}=\lfloor y\rfloor,\qquad y_{\text{frac}}=y-y_{\text{int}}\in[0,1).

则:

ex=2y=2yint2yfrac.e^x =2^y =2^{y_{\text{int}}}\cdot 2^{y_{\text{frac}}}.

整数部分 2yint2^{y_{\text{int}}} 可以通过浮点数 exponent bits 处理;小数部分 2yfrac2^{y_{\text{frac}}} 用低阶多项式近似:

2yfracPn(yfrac)=i=0npiyfraci.2^{y_{\text{frac}}} \approx P_n(y_{\text{frac}}) =\sum_{i=0}^{n}p_i y_{\text{frac}}^i.

这不是把所有 exp 都替换成多项式,而是 partial emulation:一部分走 FMA 多项式近似,另一部分仍走 MUFU,避免寄存器压力和额外指令开销抵消收益。

其他优化#

【这一部分内容待完善~】

Source Code#

【这一部分内容待完善~】

参考资料#

文章分享

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

评论区

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

音乐

暂未播放

0:000:00
暂无歌词
站点统计
文章
123
分类
20
标签
174
总字数
1,103,045
运行时长
0
最后活动
0 天前

文章目录