/ LLM算法
[LLM算法-4] FlashAttention 演进:从 IO 到 GPU 流水线
以 Memory IO、GPU 并行度和 Hopper 硬件流水线为主线,解释 FlashAttention v1、v2、v3 如何在不改变 exact attention 数学形式的前提下逐代提升性能。
FlashAttention 的三代实现可以看成一条很清晰的性能优化路线:
Memory IO -> GPU Parallelism -> Hardware Pipeline三代都在计算同一个 exact attention:
变化的不是模型结构,而是这张计算图如何映射到 GPU。v1 先解决中间矩阵的显存读写,v2 重新安排 thread block 和 warp 的工作,v3 则针对 Hopper 的 TMA、WGMMA 和 FP8 重新设计异步流水线。
1. 先问一个问题:Attention 到底慢在哪里#
标准实现通常分成三步:
S = QK^T
P = softmax(S)
O = PV当 时, 和 都是 。序列长度一大,这两个中间矩阵不仅占空间,还需要在 HBM、L2、shared memory 之间反复搬运。
因此 Attention 的瓶颈不一定是 FLOPs。现代 GPU 的 Tensor Core 能很快完成矩阵乘法,但显存带宽和片上存储容量增长得没有那么快。换句话说,很多实现首先是 memory-bound,而不是 compute-bound。
FlashAttention 的共同目标就是:让数据尽可能留在更快、更小的片上存储中,并在数据还在片上时完成更多计算。
2. FlashAttention v1:用 tiling 消灭 中间矩阵#
2.1 不再 materialize attention matrix#
FlashAttention v1(2022)的关键观察是:没有必要把完整的 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可利用的并行度从近似 扩展到:
这对 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 可以概括为:
论文在 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目标是让:
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: 算完后,必须先完成 softmax,才能继续 。而 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)也就是:
当一个操作本身很难再加速时,把它放到另一个操作的执行窗口中,是比单独优化它更有效的 latency-hiding 思路。
6. FP8:吞吐更高,也更考验数值稳定性#
Hopper 的 FP8 Tensor Core 为 Attention 提供了更高的理论吞吐,但 Attention 对量化误差敏感,尤其是 后还要经过指数运算,局部误差可能被放大。
v3 的思路包括:
- Block quantization:按 block 而不是整个 tensor 使用 scale,适应不同区域的动态范围;
- Incoherent processing:通过变换数据分布降低少数异常值对量化的影响;
- 显式的累加和缩放策略:在保持 FP8 输入吞吐的同时,控制 softmax 统计量和输出的误差。
论文报告其 FP8 Attention 的误差相较基线有所降低,并在 H100 上达到接近 1.2 PFLOPS 的峰值;这些结果同样依赖具体形状、精度和实现版本,不能简单外推到所有模型。
7. 三代 FlashAttention 的真正区别#
| FlashAttention v1 | FlashAttention v2 | FlashAttention v3 | |
|---|---|---|---|
| 主要问题 | HBM IO 太多 | GPU 并行度不足 | Hopper 硬件没有吃满 |
| 核心方法 | tiling + online softmax | sequence parallelism + warp partitioning | TMA/WGMMA + 异步流水线 |
| Thread block | 基础 tile 映射 | 跨 sequence 拆分 head | 为 pipeline 服务的角色分工 |
| Warp | 通信和同步较多 | 更独立的 Q 分工 | producer / consumer specialization |
| 数据移动 | 以同步搬运为主 | 更高效地复用 tile | TMA 异步搬运 |
| 矩阵计算 | MMA | 更高 Tensor Core 利用率 | WGMMA |
| Softmax | online softmax | 减少额外 FLOPs | 与 GEMM 交错执行 |
| 精度重点 | FP16/BF16 | FP16/BF16 | FP16/BF16 + FP8 |
可以把这条路线压缩成三个词:
8. 从 Roofline 视角串起来#
这三代也对应 Roofline Model 中瓶颈的移动:
- 原始 Attention 需要频繁读写 矩阵,Arithmetic Intensity 较低,更偏 memory-bound。
- v1 减少 HBM traffic,把计算推向更高的 Arithmetic Intensity。
- v2 发现新的限制变成 occupancy、warp communication 和 Tensor Core utilization。
- H100 的 Tensor Core 吞吐进一步提高后,数据搬运、softmax dependency 和 pipeline bubble 再次显现。
- 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 时,可以依次问四个问题:
- 它如何切分 tile,哪些数据会留在 registers 或 shared memory?
- thread block 如何覆盖 batch、head 和 sequence?
- warp 如何分工,是否存在不必要的 shared-memory exchange?
- 数据搬运、Tensor Core、softmax 是否能够重叠?
如果先建立这条从 tile 到 thread block、warp、MMA/WGMMA 的链路,看到一段 CUDA 或 Triton 代码时,就能判断它主要在优化 IO、occupancy、Tensor Core utilization,还是 latency hiding。