Back to writing

/ LLM Algorithms

[LLM Algorithms-4] The Evolution of FlashAttention: From IO to GPU Pipelines

A hardware-aware tour of FlashAttention v1, v2, and v3: how IO awareness, GPU parallelism, and Hopper pipelines accelerate exact attention without changing its mathematics.

4 minLLM · Attention · FlashAttention · CUDA

FlashAttention 的三代实现可以看成一条很清晰的性能优化路线:

Memory IO  ->  GPU Parallelism  ->  Hardware Pipeline

三代都在计算同一个 exact attention:

O=softmax(QKT)VO=\operatorname{softmax}(QK^T)V

变化的不是模型结构,而是这张计算图如何映射到 GPU。v1 先解决中间矩阵的显存读写,v2 重新安排 thread block 和 warp 的工作,v3 则针对 Hopper 的 TMA、WGMMA 和 FP8 重新设计异步流水线。

1. 先问一个问题:Attention 到底慢在哪里#

标准实现通常分成三步:

S = QK^T
P = softmax(S)
O = PV

Q,KRN×dQ,K\in\mathbb{R}^{N\times d} 时,SSPP 都是 N×NN\times N。序列长度一大,这两个中间矩阵不仅占空间,还需要在 HBM、L2、shared memory 之间反复搬运。

因此 Attention 的瓶颈不一定是 FLOPs。现代 GPU 的 Tensor Core 能很快完成矩阵乘法,但显存带宽和片上存储容量增长得没有那么快。换句话说,很多实现首先是 memory-bound,而不是 compute-bound

FlashAttention 的共同目标就是:让数据尽可能留在更快、更小的片上存储中,并在数据还在片上时完成更多计算。

2. FlashAttention v1:用 tiling 消灭 N2N^2 中间矩阵#

2.1 不再 materialize attention matrix#

FlashAttention v1(2022)的关键观察是:没有必要把完整的 N×NN\times N attention score 写回 HBM。

它把 Q、K、V 切成 tile,在 SRAM、shared memory 和 registers 中完成局部计算:

HBM
  |
  v
Q_i, K_j, V_j  ->  Q_i K_j^T  ->  online softmax  ->  P_ij V_j
                                                        |
                                                        v
                                                      O_i

一个 tile 处理完就丢弃 score,只保留输出累加所需的状态。这样就避免了:

HBM <- N x N attention matrix -> HBM

这就是 IO-aware exact attention:数学结果仍然与标准 Attention 一致,但 HBM access 显著减少,中间内存从 quadratic footprint 降下来。

2.2 Online softmax 为什么可行#

softmax 看起来需要一次看到整行分数,但分块时可以在线维护每行的最大值和归一化因子。处理新 tile 时,重新缩放之前的累加结果,再把当前 tile 的贡献合并进来。

因此可以在不保存完整 score matrix 的情况下得到与完整 softmax 等价的结果。v1 的核心不是某个单独的 CUDA 技巧,而是把数学计算重排到片上存储容量允许的范围内。

2.3 v1 的新瓶颈#

v1 减少了 HBM IO,却没有充分利用 GPU 的全部计算资源。与纯 GEMM 相比,Attention kernel 还包含 max、exp、rescale、mask 等 non-matmul 操作;同时,thread block 和 warp 的分工仍然不够理想。

结果是:GPU 可能已经少搬了很多数据,但 Tensor Core、SM 和 warp 仍然没有跑满。下一代优化的重点自然从“搬得少”转向“算得更满”。

3. FlashAttention v2:让 GPU 并行度和 Tensor Core 利用率上来#

FlashAttention v2(2023)没有改变 v1 的 IO-aware 基础,而是重新设计执行映射,重点解决三个问题。

3.1 减少 non-matmul FLOPs#

Tensor Core 擅长矩阵乘法,softmax 的 max、exp、累加和 rescale 则主要由普通 CUDA core 完成。对 Attention 来说,这些非矩阵运算的比例并不小,而且会引入依赖和同步。

v2 调整 online softmax 的更新方式,减少不必要的 rescaling 和其他 non-matmul FLOPs,让总执行时间更接近两次主要 GEMM:

QK^T  +  PV  >>  其他逐元素操作

这里的目标不是让 exp 突然变快,而是减少它在整个 kernel 中占据的比例。

3.2 Sequence parallelism:一个 head 不再只有一个 block#

如果并行只沿 batch 和 head 展开,那么小 batch、少 head 或长序列场景下,GPU 可能只有很少的主要任务。例如一个 batch 有 16 个 head,而 GPU 有上百个 SM,很多 SM 会处于空闲状态。

v2 允许同一个 Attention head 沿 sequence 维度拆给多个 thread block:

Head 0:  [ TB0 | TB1 | TB2 | TB3 | ... ]
          sequence tiles

可利用的并行度从近似 B×HB\times H 扩展到:

B×H×sequence tilesB\times H\times\text{sequence tiles}

这对 long context、small batch 的推理尤其重要。需要注意的是,sequence parallelism 会带来部分归约和同步成本,所以它不是无条件地越细越好;kernel 需要在 occupancy、局部性和归约开销之间做平衡。

3.3 改进 warp work partitioning#

在较早的分工方式中,多个 warp 可能共同处理 K/V,并通过 shared memory 交换中间结果:

warp 0 --+
warp 1 --+--> shared memory --> combine
warp 2 --+

虽然 shared memory 比 HBM 快,但仍然比 registers 慢,而且会引入同步。v2 让不同 warp 更独立地处理 Q 的不同部分,同时复用相同的 K/V tile:

                 K, V tile
              /      |      \\
           warp 0  warp 1  warp 2
             Q0      Q1      Q2

这样减少 warp 之间的 shared-memory communication,让更多时间真正用于 Tensor Core 计算。

3.4 v2 的位置#

因此,v2 可以概括为:

less non-matmul work+more parallelism+less warp communication\boxed{\text{less non-matmul work} + \text{more parallelism} + \text{less warp communication}}

论文在 A100 等配置上报告了相较 v1 的显著加速,并将利用率从约 25%--40% 提升到约 50%--73% 的理论 FLOPs。具体数字会随序列长度、head dimension、数据类型和 kernel 配置变化,不能脱离测试条件直接当成固定保证。

4. Hopper 带来新硬件:v2 又暴露出新瓶颈#

H100(Hopper)提供了 Ampere 上没有的关键能力:

硬件能力对 Attention kernel 的意义
TMA(Tensor Memory Accelerator)异步地把多维 tensor 从 HBM 搬到 shared memory
WGMMA以 warpgroup 为单位使用 Tensor Core 做矩阵乘加
FP8 Tensor Core用更低精度换取更高吞吐,但需要额外数值控制
异步执行机制让数据移动、矩阵计算和部分标量计算重叠

如果只是把偏向 Ampere 的 v2 kernel 原样搬到 H100,仍然可能按“搬完再算、算完再搬”的同步节奏运行,无法发挥 Hopper 的硬件吞吐。此时,数据移动和 softmax 的等待又成为新的瓶颈。

5. FlashAttention v3:用异步流水线隐藏延迟#

FlashAttention v3(2024)的中心思想可以用一个词概括:asynchrony

5.1 TMA load 与 WGMMA compute 重叠#

同步式执行大致是:

load tile 0 -> compute tile 0 -> load tile 1 -> compute tile 1 -> ...

当 Tensor Core 在计算时,内存流水线可能空闲;当内存加载时,Tensor Core 也可能空闲。v3 使用 TMA 和 Hopper 的异步机制,把它改造成流水线:

时间 --->
 
TMA:       load 0       load 1       load 2       load 3
WGMMA:          compute 0     compute 1     compute 2     compute 3

目标是让:

memory transferTensor Core compute\text{memory transfer}\parallel\text{Tensor Core compute}

5.2 Warp specialization:producer 和 consumer#

v3 进一步明确 warp 的角色:

Producer warp      -> TMA load,HBM -> shared memory
Consumer warpgroup -> WGMMA,执行 QK 和 PV

这类似一个生产者--消费者流水线:负责搬数据的 warp 不必等待计算完成,负责矩阵乘法的 warp 也不必亲自处理每次数据搬运。关键收益不是减少某一条指令的延迟,而是把原本暴露在关键路径上的等待藏到其他工作之后。

5.3 让 GEMM 和 softmax 互相遮蔽#

Attention 还有一个天然的 pipeline bubble:QKTQK^T 算完后,必须先完成 softmax,才能继续 PVPV。而 softmax 的 max、exp、sum 并不能直接使用 Tensor Core。

v3 采用 block-wise interleaving,让不同 tile 的工作交错进行:

Tensor Core:  QK(tile i+1)   QK(tile i+2)
CUDA core:        softmax(tile i)   softmax(tile i+1)

也就是:

GEMMi+1Softmaxi\boxed{\text{GEMM}_{i+1}\parallel\text{Softmax}_i}

当一个操作本身很难再加速时,把它放到另一个操作的执行窗口中,是比单独优化它更有效的 latency-hiding 思路。

6. FP8:吞吐更高,也更考验数值稳定性#

Hopper 的 FP8 Tensor Core 为 Attention 提供了更高的理论吞吐,但 Attention 对量化误差敏感,尤其是 QKTQK^T 后还要经过指数运算,局部误差可能被放大。

v3 的思路包括:

  • Block quantization:按 block 而不是整个 tensor 使用 scale,适应不同区域的动态范围;
  • Incoherent processing:通过变换数据分布降低少数异常值对量化的影响;
  • 显式的累加和缩放策略:在保持 FP8 输入吞吐的同时,控制 softmax 统计量和输出的误差。

论文报告其 FP8 Attention 的误差相较基线有所降低,并在 H100 上达到接近 1.2 PFLOPS 的峰值;这些结果同样依赖具体形状、精度和实现版本,不能简单外推到所有模型。

7. 三代 FlashAttention 的真正区别#

FlashAttention v1FlashAttention v2FlashAttention v3
主要问题HBM IO 太多GPU 并行度不足Hopper 硬件没有吃满
核心方法tiling + online softmaxsequence parallelism + warp partitioningTMA/WGMMA + 异步流水线
Thread block基础 tile 映射跨 sequence 拆分 head为 pipeline 服务的角色分工
Warp通信和同步较多更独立的 Q 分工producer / consumer specialization
数据移动以同步搬运为主更高效地复用 tileTMA 异步搬运
矩阵计算MMA更高 Tensor Core 利用率WGMMA
Softmaxonline softmax减少额外 FLOPs与 GEMM 交错执行
精度重点FP16/BF16FP16/BF16FP16/BF16 + FP8

可以把这条路线压缩成三个词:

IO awarenessparallelism awarenesshardware/pipeline awareness\boxed{\text{IO awareness}\rightarrow\text{parallelism awareness}\rightarrow\text{hardware/pipeline awareness}}

8. 从 Roofline 视角串起来#

这三代也对应 Roofline Model 中瓶颈的移动:

  1. 原始 Attention 需要频繁读写 N×NN\times N 矩阵,Arithmetic Intensity 较低,更偏 memory-bound。
  2. v1 减少 HBM traffic,把计算推向更高的 Arithmetic Intensity。
  3. v2 发现新的限制变成 occupancy、warp communication 和 Tensor Core utilization。
  4. H100 的 Tensor Core 吞吐进一步提高后,数据搬运、softmax dependency 和 pipeline bubble 再次显现。
  5. v3 用异步流水线把这些延迟藏在矩阵计算后面,尽量让 GPU 接近 compute-bound。

这体现了一个通用的性能规律:

优化掉一个瓶颈以后,下一个瓶颈才会暴露出来。

9. 给推理 kernel 的一个阅读框架#

不要把 v1、v2、v3 只当成 Attention 的“算法版本”。更准确的抽象是:

Attention 数学公式(基本不变)
              |
              v
       kernel mapping
        /           \\
   memory mapping  compute mapping
        \\           /
              v
        GPU architecture

阅读 vLLM、FlashInfer 或 Triton Attention kernel 时,可以依次问四个问题:

  1. 它如何切分 tile,哪些数据会留在 registers 或 shared memory?
  2. thread block 如何覆盖 batch、head 和 sequence?
  3. warp 如何分工,是否存在不必要的 shared-memory exchange?
  4. 数据搬运、Tensor Core、softmax 是否能够重叠?

如果先建立这条从 tile 到 thread block、warp、MMA/WGMMA 的链路,看到一段 CUDA 或 Triton 代码时,就能判断它主要在优化 IO、occupancy、Tensor Core utilization,还是 latency hiding。

参考资料#