
FlashAttention 源码级深度解析:从 IO 感知 Tiling 与 Online Softmax 到 Hopper/Blackwell 异步流水线的注意力内核底层原理核心痛点:Transformer 的注意力计算显存与带宽开销随序列长度平方增长,标准实现被 HBM 带宽墙卡死,无法支撑长上下文大模型的高吞吐训练与推理适配人群:具备 CUDA/C++ 基础、希望深入算子内核源码的 AI 工程师、GPU 内核开发者、推理引擎(vLLM/SGLang/TensorRT-LLM)优化者收获能力:掌握注意力内核「IO 感知 tiling + online softmax + 反向重计算」三大核心算法,理解 FlashAttention-1/2/3/4 四代在 Ampere/Hopper/Blackwell 硬件上如何逐代突破瓶颈,并能独立解读内核源码与落地部署技术背景与演进逻辑注意力是 Transformer 的算力与显存黑洞:自注意力要求每个 query 与全部 key 做内积,得分矩阵 S 与概率矩阵 P 均为 N×N 规模,N 为序列长度长上下文直接放大这个平方项:N=4096 时 P 矩阵已 4096×4096 浮点 = 64MB(fp32),N=128k 时单个头就需 64GB,远超任何单卡 HBM 容量标准实现(如早期 PyTorch 的 softmax(QK^T)V)把 S、P 完整物化到 HBM,反复读写,导致注意力几乎完全被内存带宽而非算力约束业界戏称注意力是「memory-bound 到发指」:GPU 的 HBM 带宽(约 3TB/s)远跟不上 Tensor Core 的算力(约 2000 TFLOPS),两者增速长期失衡关键洞察(IO 感知):注意力并不需要把整个 S、P 写回 HBM——只要把数据切成能装进片上 SRAM 的小块(tile),在 SRAM 内完成 softmax 并累加,就能把 HBM 读写量从 O(N²) 级降到 O(N) 级这正是 FlashAttention 的核心思想,也是它与「单纯工程优化」的本质区别:它在算法层面改变了访存复杂度,而不只是把同一算法跑得更快演进逻辑:每一代 FlashAttention 都对应「算法侧 + 硬件侧」的一次协同突破,从不新造硬件,而是把当时代 GPU 的新特性用满演进时间线(text 树表达):演进时间线 ├── 2022 FlashAttention-1 - IO 感知 tiling + online softmax + 反向重计算,Ampere A100,2-4x 加速 ├── 2023 FlashAttention-2 - 减少非矩阵乘 FLOPs + 序列维度并行 + warp 分工,2x 再加速,A100/H100 ├── 2024 FlashAttention-3 - Hopper 异步(TMA/WGMMA) + warp specialization + FP8,H100 利用率 35%-75% └── 2026 FlashAttention-4 - Blackwell 非对称扩展 + TMEM + 2-CTA MMA + 软件 exp2,B200 达 1613 TFLOPS/s总结:FlashAttention 家族的演进逻辑是「算力与带宽增速失衡」这一硬件事实催生的必然结果——当 Tensor Core 每代翻倍、而 SFU 与共享内存带宽原地踏步时,注意力内核的瓶颈会从「带宽」漂移到「指数运算」再到「共享内存流量」,每一代都在针对新的最短板重写流水线核心原理深度解析原理模块一:IO 感知 tiling(FlashAttention-1 的灵魂)标准注意力的访存灾难标准注意力三步走:S = Q K T S = Q K^TS=QKT,P = m a t h r m s o f t m a x ( S ) P = mathrm{softmax}(S)P=mathrmsoftmax(S),O = P V O = P VO=PV,其中 Q/K/V/O 均为 N×d 矩阵每一步都把中间结果写回 HBM 再读回:先写 S(N×N),再写 P(N×N),整个流程的 HBM 访问量是 O(N²)公式(IO 复杂度):标准实现 HBM 访问 ≈O ( N d + N 2 ) O(N d + N^2)O(Nd+N2),而片上 SRAM 只有约 200KB,N×N 矩阵根本放不下Tiling 的数学本质:softmax 可分解softmax 的归一化因子(分母)可以分块累加,这是 tiling 得以成立的数学根基把 K、V 沿序列维度切成 B_r 块,Q 切成 B_c 块,每块 Q 与一块 K 算出局部 S 子块,立即在 SRAM 内算局部 softmax 并累加进输出,S/P 从不写回 HBM公式(IO 复杂度下降):tiling 后 HBM 访问 ≈O ( N 2 d 2 / M ) O(N^2 d^2 / M)O(N2d2/M),其中 M 为 SRAM 容量,d 为头维度——当 d=64、M≈100KB 时,访问量下降一个数量级设计思想核心洞察:既然数据从 HBM 读到 SRAM 的代价远高于 SRAM 内的计算,那就让数据「一次读入、多次复用」,把对 HBM 的访问次数压缩到理论下界原理模块二:online softmax(数值稳定 + 分块可累加)为什么需要 online softmax朴素 softmax 需要先遍历整行求最大值 m 做数值稳定(避免e x e^xex上溢),这要求把整行 S 读两遍;分块处理时无法预先知道全局最大值online softmax 用「运行最大值 + 运行和」的增量更新,一次遍历即可得到与全量 softmax 数值等价的结果增量更新公式维护运行最大值 m、运行和 l、运行输出 O,处理新块 j 时:m n e w = m a x ( m o l d , m j ) m_{new} = max(m_{old}, m_j)mnew=max(mold,mj