FlashAttention 并没有把 dense attention 从
| 方法 | 核心问题 | 主要优化层次 |
|---|---|---|
| Standard Attention | 直接按算子执行,产生大量 |
|
| FlashAttention | 同样计算 exact attention,如何显著减少 HBM 数据搬运 | HBM |
| FlashAttention-2 | IO 已经优化后,如何进一步提高 GPU 利用率 | Thread Block / Warp / Tensor Core |
因此,理解 FlashAttention 系列最重要的视角不是“减少了多少 FLOPs”,而是:
graph LR A(Standard Attention) --> B(IO-aware Attention) --> C(GPU-efficient IO-aware Attention)
Attention 的性能瓶颈
设单个 attention head 的
标准 Attention 可以写成:
其中
从算法复杂度看,两次主要矩阵乘都需要
标准实现通常经历:
也就是说,完整的
以 A100 为例,HBM 容量很大、带宽约为 TB/s 量级,而每个 SM 的片上 SRAM 容量只有数百 KB,却拥有远高于 HBM 的带宽。于是 GPU kernel 大致可分为两类:大型 GEMM 往往具有较高 arithmetic intensity,更接近 compute-bound;softmax、mask、dropout、reduction 等操作 FLOPs 很少,但需要反复读取和写回数据,更容易 memory-bound。
对 memory-bound workload,更合理的性能模型是:
而不是只看 FLOPs。FlashAttention 的核心出发点正是:Attention
的主要瓶颈之一,是
FlashAttention
FlashAttention 的核心目标是
它把
真正困难的地方不是矩阵乘,而是 softmax。矩阵乘天然可以分块:
但一行 softmax 的归一化依赖整行所有元素。FlashAttention 能够分块计算 exact softmax,依赖的就是 online softmax。
Online Softmax
考虑一行 score:
为了数值稳定,softmax 通常写成:
如果 SRAM 中一次只能看到一个 score
block,我们暂时不知道最终全局最大值
假设已经处理过一部分 scores,保存:
新 block
新的全局最大值:
由于
旧 denominator 可以直接重标定:
因此,完整 softmax 不需要保存完整 score row,只需要持续维护 row-wise
的
FlashAttention 前向过程
论文 Algorithm 1 的输入是:
它们存放在 HBM 中,同时假设片上 SRAM 可用容量为
随后把:
其中
并有:
算法为每个 query row 维护三个状态:
初始化为:
其中
第一版 FlashAttention 的循环顺序是:
1 | for each K_j, V_j block: |
外层固定一个
对于当前 tile,首先计算:
实际 Transformer 中还包含
然后对每一行计算当前 tile 的局部统计量:
这里的
接下来将旧状态与当前 block 合并。新的最大值:
新的 denominator:
这一步就是 online softmax 的 row-wise recurrence。
仅更新 denominator 还不够,因为 Attention 最终需要的是
如果旧结果已经是归一化后的
当最大值从
论文伪代码中的 diag 只是把每一行各自的
这一递推始终维持三个不变量。处理到第
当所有 K/V blocks 处理完成后,
这也解释了为什么 FlashAttention 可以“边算边丢”。一个局部 tile 的生命周期是:
完整的
Tiling、Fusion 与 Backward Recomputation
FlashAttention 的 tiling 与 kernel fusion 是同一件事的两个侧面。标准 Attention 可能跨多个 kernel 依次执行:
FlashAttention 则把:
尽可能融合在同一个 kernel 的片上执行路径中。online softmax 使 tile
可以独立推进,tiling 又使
训练时 backward 更能体现“IO 优先”的设计思路。标准 Attention backward 涉及:
以及:
最直接的实现会在 forward 保存完整
再利用 forward 保存的 normalization statistics 重建:
因此 FlashAttention 的一个关键工程原则是:
也就是:如果重新计算主要发生在 Tensor Core、register 和 SRAM 中,而它替代的是大量 HBM 读写,那么增加 FLOPs 反而可能缩短 wall-clock time。
Softmax backward 还有一个重要化简。定义:
由于:
可以得到:
于是局部 backward 可写成:
其中
复杂度
FlashAttention 没有改变 dense attention 的主计算复杂度:
它改变的是中间存储和 HBM IO。完整的
论文给出的 HBM IO complexity 对比为:
其中
论文中的一组 GPT-2 medium 实验
| 指标 | Standard Attention | FlashAttention |
|---|---|---|
| GFLOPs | 66.6 | 75.2 |
| HBM R/W | 40.3 GB | 4.4 GB |
| Forward + Backward | 41.7 ms | 7.3 ms |
FlashAttention 的 FLOPs 甚至更多,但 HBM 读写从 40.3 GB 降到 4.4 GB,运行时间反而从 41.7 ms 降到 7.3 ms。这个结果很好地说明了 FlashAttention 的核心 它主要降低的是 IO complexity,而不是主 FLOP complexity.
FlashAttention-2
FlashAttention 已经解决了最严重的 HBM IO 问题,但第一版 kernel 的吞吐仍明显低于高质量 GEMM。FlashAttention-2 的 profiling 发现,新的瓶颈集中在 GPU work partitioning:thread block 数量不足、occupancy 不够、warp 分工不理想、shared-memory communication 偏多,以及 non-matmul FLOPs 的实际代价较高。
因此 FA2 的重点从:
继续下沉到:
减少 Non-Matmul FLOPs
在 A100 上,FP16/BF16 Tensor Core matmul 的理论吞吐远高于普通 FP32 non-matmul operation。也就是说,从硬件执行成本看,一个 elementwise/reduction FLOP 与一个 Tensor Core matmul FLOP 并不“等价”。
FA1 在每个 K/V tile 后维护已经归一化的
当新 block 到来时:
只有所有 column blocks 处理完成后,才做一次:
数学结果不变,但 inner loop 中减少了昂贵的 non-matmul work。
FA2 还把 forward 保存的 softmax statistics 压缩成 LogSumExp。由于:
定义:
即可写成:
因此 backward 不必同时依赖
Sequence Parallelism
第一版 FlashAttention 的主要并行维度近似是:
如果 batch 很小、head 数也有限,就可能没有足够的 thread blocks 填满 GPU。长序列尤其容易出现这个问题,因为显存压力往往迫使 batch size 进一步减小。
FA2 将 Q 的 row block 也纳入并行维度。FA1 的逻辑顺序近似是:
1 | for K/V column block: |
FA2 forward 改为:
1 | parallel for Q row block: |
也就是把:
作为 thread-block 级的独立工作单元。并行任务数量由近似:
提升到:
其中:
例如:
则:
并行 work units 从约
这使大量 SM 更容易保持忙碌。
Forward 适合 row-oriented parallelism,因为每个
Warp Work Partitioning
增加 thread-block 数量解决的是 block-level parallelism,FA2 还进一步优化一个 block 内多个 warp 的分工方式。
FA1 forward 的典型方案可以理解为 Split-K:多个 warp
共享同一组 Q rows,而把 K/V 方向上的工作拆开。这样在
于是需要把 partial results 写入 shared memory、同步、重新读取并 reduction。真正昂贵的不是几次加法,而是:
FA2 改为 Split-Q。把 Q rows 在 warp 之间切分:
而 K/V 对这些 warp 可复用。每个 warp 负责独立的一部分 output rows:
最终输出是拼接:
而不是多个 partial output 的求和,因此不需要 warp 间 reduction。
这体现了一个很通用的 GPU kernel 原则:
后者通常意味着更多 synchronization、shared-memory traffic 和 reduction。
Tile Size 与 Causal Attention
FlashAttention 并不是 tile 越大越好。增大
因此实际调优必须同时考虑:
Causal Attention 还可以利用 tile 结构进一步减少无效计算。对于:
如果一个 tile 完全位于 causal mask 的上三角区域,那么整个 tile
都不会贡献输出,可以直接跳过,不必计算
| 维度 | FlashAttention | FlashAttention-2 |
|---|---|---|
| Attention 数学定义 | Exact dense attention | Exact dense attention |
| 主计算复杂度 | ||
| 主要瓶颈 | HBM IO | GPU utilization |
| 核心优化层次 | HBM |
Thread Block / Warp |
| Softmax | Online softmax | Online softmax + 算术精简 |
| 不写入 HBM | 不写入 HBM | |
| Forward accumulator | 每 block 更新 normalized |
维护 unnormalized |
| Forward statistics | LogSumExp |
|
| Backward | Recomputation | Recomputation + 更好的 work partition |
| 主要并行维度 | Batch × Heads | Batch × Heads × Sequence |
| Forward loop orientation | K/V → Q | Q → K/V |
| Warp partition | Split-K | Split-Q |
| Shared-memory communication | 较多 | 更少 |
从 Roofline Model 看,两代算法的关系尤其清楚。一个 kernel 的实际性能受:
限制,其中:
是 arithmetic intensity。
标准 Attention 中,大量
FA1 的成功并没有结束优化,反而暴露了 FA2 要解决的问题。
两篇论文最值得迁移到其他高性能 kernel 的经验也可以归纳为同一条逻辑:算法复杂度不能替代硬件性能模型;数据移动本身就是成本;recomputation 不一定比 memory access 贵;并行任务数量和 work partitioning 同样重要;tile size 必须在数据复用、寄存器、shared memory 与 occupancy 之间综合权衡。
总结
FlashAttention 的突破,不是发明了新的 Attention
数学形式,而是重新安排 exact attention 的执行顺序。通过 tiling、online
softmax、kernel fusion 和 backward recomputation,它避免把完整的
FlashAttention-2 则进一步认识到,当 IO 瓶颈被缓解后,性能会转而受 GPU 并行映射限制。于是它减少 inner loop 中的 non-matmul FLOPs,引入 sequence-level thread-block parallelism,并将 warp work partition 从 Split-K 改成 Split-Q,减少 shared-memory communication,提高 occupancy 和 Tensor Core utilization。
因此,FlashAttention 系列真正展示的是一条完整的 algorithm-hardware co-design 路径:
高性能 Attention 的关键不只是“算多少”,而是同时设计:数据放在哪里、什么时候搬、以什么粒度计算、由谁计算,以及中间结果是否值得保存。
讨论
评论