Flashattention

Flashattention

FlashAttention 入门:从 FA1 到 FA4 的 IO 感知注意力

FlashAttention 入门:从 FA1 到 FA4 的 IO 感知注意力 注意力层是把 Transformer 扩展到长序列时的主要瓶颈:它的时间与内存开销都随序列长度平方增长。FlashAttention 是一族让注意力保持精确(exact)的同时显著提速、省内存的算法,其核心理念是把 GPU 的内存流量——而不只是 FLOPs——当作需要优化的资源。本文用原创的图示与推导讲清核心思想:GPU 内存层级、分块(tiling)、在线 softmax(online softmax)、反向重计算(recomputation),然后是 FA1 → FA2 → FA3 → FA4 的时间线、独立 flash-attn 包与 PyTorch scaled_dot_product_attention 的区别,以及实践中真正要紧的硬件注意事项。 本文有两条规则。第一,所有性能数字均为引用,来自原始论文、Dao-AILab 官方仓库与 PyTorch 官方文档,每个数字都绑定其测量硬件;本文没有任何本地基准结果。第二,本文写作所用的工作站 GPU 是 TITAN RTX(Turing 架构,计算能力 7.5)——这不是巧合,因为官方的 FlashAttention-2/3/4 CUDA 内核无法在 Turing 上运行,最后一节把这个限制转化为一份诚实的验证方案。 文中所有”截至”表述均指 2026-08-09。 注意力为什么是平方复杂度 以单头为例,设 query、key、value 矩阵为 $Q, K, V \in \mathbb{R}^{N \times d}$($N$ 个 token,头维度 $d$),缩放点积注意力为: $$S = \frac{Q K^{\top}}{\sqrt{d}}, \qquad P = \operatorname{softmax}(S) \text{(按行)}, \qquad O = P V,$$

字体
阴影
滤镜
圆角
主题色