关于FlashAttention的一些思考
datawhale社区:Datawhale-学用 AI,从此开始
llm-algo-leetcode教程地址:GitHub - datawhalechina/llm-algo-leetcode: LLM algorithm practice lab with theory, solutions, and test cases.《大模型算法与系统教程》面向大模型入门到进阶的算法实战教程,覆盖原理讲解、答案解析、测试用例与 CUDA/Triton 实战。 · GitHub
- FlashAttention 的 FLOPs 其实不减反增(反向传播时它不保存 N×N 的注意力矩阵,而是用 Recomputation 重算一遍),但它仍然大幅提速——这恰好证明了瓶颈在访存而非计算。它没减少计算量,甚至略有增加,快是因为砍掉了 HBM 读写。
- FlashAttention 的演进不是单纯的"算法越来越聪明",而是"算法与 GPU 硬件一代代互相磨合、协同设计(hardware-algorithm co-design)"的过程。
演进的主轴
版本 解决的核心问题 优化层次 与硬件的关系 V1 显存墙:中间矩阵不能落显存 数学层(online softmax + tiling) 硬件无关,通用 V2 GPU 吃不饱:慢操作多、并行度低 算法调度层(循环重排、推迟归一化) 仍较通用 V3 榨干特定硬件:异步化一切 指令/微架构层(WGMMA、TMA、流水线) 强绑定 Hopper V4 新架构继续工程化 代码生成 + kernel 组织 面向 Blackwell FlashAttention 的演进是一条算法-硬件协同设计的曲线:V1 用数学证明注意力可以分块在线计算,解决了"能不能省显存"的问题;V2 靠重排计算顺序解决"GPU 利用率高不高"的问题;从 V3 开始,优化下沉到特定 GPU 的专属硬件单元(异步矩阵指令、硬件搬运器、流水线),进入了"为某一代芯片量身定制"的阶段。越往后,算法创新占比越小,工程与硬件绑定的占比越大。
