音乐
暂未播放
FlashAttention 1~4 全解汇总

大家好,我是芯缘,是 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) S 和 P 均为 O(N2) ,被反复读写 HBM 时,即使 Attention 的计算量为 O(N2d),整体容易呈现 Memory-Bound。
Solution#
沿序列长度分块,在 SRAM 中完成局部 score、online softmax 和输出累加,从而避免将完整的 S 和 P 写入 HBM。
Input: Output: Q,K,V∈RN×d,SRAM size=MO=softmax(QK⊤)V// Step 1: Set block sizeBc←⌈4dM⌉,Br←min(⌈4dM⌉,d)// Step 2: InitializeO←0,ℓ←0,m←−∞// Step 3: Split by sequence lengthQ→{Q1,…,QTr},Qi∈RBr×dK→{K1,…,KTc},Kj∈RBc×dV→{V1,…,VTc},Vj∈RBc×dO→{O1,…,OTr},ℓ→{ℓ1,…,ℓTr},m→{m1,…,mTr}// Step 4: Forward loopfor j←1 to Tc doKj,Vj: HBM → SRAMfor i←1 to Tr doQi,Oi,ℓi,mi: HBM → SRAMSij←QiKj⊤minew←max(mi,rowmax(Sij))Pij←exp(Sij−minew)ℓinew←emi−minewℓi+rowsum(Pij)Oinew←ℓinewemi−minewℓiOi+PijVjOi,ℓi,mi←Oinew,ℓinew,minewOi,ℓi,mi: SRAM → HBMend forend for// Step 5: Returnreturn ODetails#
Block Size设置 关键约束是一次计算时,Kj,Vj、Qi,Oi 和局部 score block Sij 都要放进 On-chip SRAM:
Bcd=O(M),Brd=O(M),BrBc=O(M).其中 Bcd=O(M) 来自 Kj,Vj 的驻留需求,Brd=O(M) 来自 Qi,Oi 的驻留需求,BrBc=O(M) 来自 Sij=QiKj⊤∈RBr×Bc 的临时存储需求。
FA1 让 Bc 尽量大:
Bc=Θ(dM),这样外层 K,V block 数更少,Q,O 被重复扫描的轮数也更少。但 Bc 变大后,Sij 的约束会限制 Br:
BrBc=O(M)⇒Br=O(d).因此论文中取:
Bc=⌈4dM⌉,Br=min(⌈4dM⌉,d).为什么 m 初始化为 −∞? m 记录的是当前已经扫描过的 score 的行最大值。online softmax 在处理第一个 Sij block 时,希望新最大值完全由这个 block 决定:
minew=max(mi,rowmax(Sij))=rowmax(Sij).把 mi 初始化成 −∞,就相当于告诉算法“当前还没有看过任何 score”。同时旧分母和旧输出的缩放项会自然消失:
emi−minewℓi=e−∞−minew⋅0=0.这样第一块和后续块可以使用完全相同的更新公式,不需要单独写 first block 的特殊分支。
复杂度分析 计算量仍然由两次矩阵乘主导,总体 FLOPs 仍为:
O(N2d).HBM 访问量的变化更关键。所有 K,V block 在外层循环中合计读入一次,代价是 Θ(Nd);每处理一个 K,V block,需要从 HBM 扫描一遍全部 Q,O block 以及 m,ℓ 统计量,主导代价是 Θ(Nd),不是 Θ(Nd+N2)。
这里不会出现每轮 Θ(N2) 的 HBM 访问,因为局部 Sij 和 Pij 虽然总体计算量覆盖了 N2 个 attention score,但它们只在 SRAM 中短暂生成和消费,并不会作为完整矩阵写入 HBM。因为:
Tc=BcN=Θ(MNd),所以 HBM 访问量为:
Θ(Nd+NdTc)=Θ(NdTc)=Θ(MN2d2).标准 Attention 需要把完整 S 和 P 写入并读出 HBM,HBM 访问量为:
Θ(Nd+N2).当片上 SRAM 足够容纳较大的 block,且通常 M≫d2 时,FlashAttention-1 可以显著减少 HBM 访问。 注:SRAM容量通常为约 64 到 164 KiB;FP16/BF16 下约 32K 到 82K 个元素。
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 操作。
Details#
通过交换内外层循环带来的收益
(1) 只需要在外循环里内循环外对 FA1 Oi进行归一化的rescale。
(2) 因为 Q 在外循环,所以不需要每次都把临时的O,m,ℓ进行FA1那样内外循环读写,而是每次外循环都是赋值0。
(3)在forward过程中维护的用于训练的参数m,ℓ不需要保存到 HBM,而是只保存 logsumexp:Li=mi+log(ℓi) 注:如果是纯推理的话,m,ℓ,L其实都不需要写回HBM,L推理中都不需要有,但是m,ℓ无论训推无论FA version,都需要中途维护。
复杂度分析
总体 FLOPs 仍为:O(N2d).
如果按 FA2 的 row-block outer loop 粗略估算 HBM 访问量,每个 Qi block 会扫描一遍完整 K,V:
Tr=BrN.Tr=BrN,HBM reads of K,V≈Θ(NdTr).HBM 访问估算:
Θ(NdTr)=Θ(BrN2d).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=emj−1−mjOj−1+eSj−mjVj.FA4 观察到,只有当新的最大值明显变大时,旧输出才必须立即 rescale。因此它区分真实候选最大值 mj 和当前用于缩放的基准 mˉj−1,只在:
mj−mˉj−1>τ时才更新缩放基准并执行缩放:
mˉj←mj,Oj=emˉj−1−mˉjOj−1+eSj−mˉjVj.否则不更新用于缩放的 mˉ,也不对旧输出做 rescale:
mˉj←mˉj−1,Oj=Oj−1+eSj−mˉj−1Vj.阈值 τ 控制可以容忍的 slack。论文中典型取 τ=log2(256)=8,对应“缩放因子不超过 256 时先不 rescale”。如果用自然指数 ex 的写法理解,对应阈值就是 ln(256)。最后再用最终的 true normalizer 统一归一化,减少中间反复缩放带来的 non-matmul 开销。
Software-Emulated Exponential#
Blackwell 上 Tensor Core 吞吐增长很快,但指数函数单元(MUFU)没有同比例增长,softmax 里的 exp 会成为瓶颈。FA4 因此用 FMA 单元分担一部分指数计算。
softmax 数学上使用 ex,但实现中可以先转成 base-2 exponent:
ex=2xlog2e.令:
y=xlog2e,yint=⌊y⌋,yfrac=y−yint∈[0,1).则:
ex=2y=2yint⋅2yfrac.整数部分 2yint 可以通过浮点数 exponent bits 处理;小数部分 2yfrac 用低阶多项式近似:
2yfrac≈Pn(yfrac)=i=0∑npiyfraci.这不是把所有 exp 都替换成多项式,而是 partial emulation:一部分走 FMA 多项式近似,另一部分仍走 MUFU,避免寄存器压力和额外指令开销抵消收益。
其他优化#
【这一部分内容待完善~】
Source Code#
【这一部分内容待完善~】
参考资料#
- Datawhale: FlashAttention Sim
- Datawhale: FlashAttention Memory Model
- FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness
- FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning
- FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision
- FlashAttention-4: Algorithm and Kernel Pipelining Co-Design for Asymmetric Hardware Scaling
- FlashAttention 完全拆解(一):IO 感知与分块注意力算法(本站)
- FlashAttention 完全拆解(二):序列维并行、Warp 工作划分与 FlashDecoding(本站)
- FlashAttention 完全拆解(三):Warp 特化异步流水线与 FP8 低精度(本站)
- FlashAttention 完全拆解(四):FlashAttention-4——面向 Blackwell 的算法-流水线协同设计(本站)
文章分享
如果这篇文章对你有帮助,欢迎分享给更多人!
部分内容可能已过时
评论区
分享你的想法,与大家交流讨论
音乐
暂未播放



