FlashAttention 算法深度解析:从 IO-Aware 优化到 GPU 硬件极致
摘要:FlashAttention 通过 IO-Aware 的 Tiling + Recomputation 策略,将 Self-Attention 的 HBM 访问复杂度从 O(N²) 降至 O(N²/d),在不牺牲数学精度的前提下实现 2-4× 端到端训练加速。本文深入推导 Online Softmax 数学基础、解析 FlashAttention/2/3 三代算法演进、剖析 Hopper/Tensor Core 适配优化,并给出一套完整的工程性能分析框架。
一、引言:Attention 的 IO 墙困境
2017 年 Transformer 架构诞生以来,Self-Attention 机制已成为大语言模型的核心算子。其数学表达简洁优美:
\[Attention(Q, K, V) = \text{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right) V$\]
但看似简洁的背后隐藏着严峻的硬件效率问题。设序列长度 N、特征维度 d,标准 Attention 的三步计算流程为:
- S = QK^T:计算 N×N 注意力分数矩阵,写入 HBM
- P = softmax(S):HBM 读取 S,计算 softmax,结果写入 HBM
- O = PV:HBM 读取 P 和 V,矩阵乘法得到输出
- 算力:FP16 Tensor Core 312 TFLOPS
- HBM 带宽:2 TB/s
- 计算/带宽比 ≈ 156
- 子块 1:最大值 $m^{(1)}$,归一化分母 $\ell^{(1)} = \sum e^{x_i^{(1)}-m^{(1)}}$
- 子块 2:最大值 $m^{(2)}$,归一化分母 $\ell^{(2)} = \sum e^{x_j^{(2)}-m^{(2)}}$
- 标准 Attention:$QK^T$ (O(N²) 写) + softmax (O(N²) 读/写) + $PV$ (O(N²) 读)
- 总计:$O(N^2)$ HBM 访问
- FlashAttention:对负载入逐步块的 K, V,每个 token 的 Q 只需计算一次 Attn,但块的 K, V 会多个 Q 块多次载入
- 总计:$O(N^2 d^2 / M)$ HBM 访问,其中 M 为 SRAM 大小
- 当 N >> d 时,访问次数大幅减少
- FlashAttention-1:使用 CUDA Thread Block (warps × 1),单个 block 处理连续的行范围
- FlashAttention-2:单个 Thread Block 内多个 warp 处理不同行,共享 K, V 块的数据但独立计算 softmax
- 长序列训练(N > 2048):IO 节省效果显著
- Batch size 受限场景:减少显存占用允许更大 batch
- 大模型预训练:Attention 超时占比高(30-50%)
- 推理短序列(N < 512):Tiling 开销占比可能反而不划算
- 已经使用线性注意力或其他 O(N) 退化的模型
- 需要精确保留 N×N Attention 权重的可解释性分析
- 因果掩码 (Causal Mask) 实现:FL2 通过设置 j < i + 1 的掩码矩阵避免未来信息泄漏,但跳过整个 j < i 的分块计算可以节省约 50% 计算量。
- 变长序列 (variable length):FL2 支持 varlen 接口,通过 cu_seqlens 参数将不同样本打包到一个 batch 中处理,避免 padding 浪费。
- Dropout 实现:需要在 SRAM 中生成随机掩码,但只在非块对角线元素上计算才能保持正确性。
- Alibi 位置编码:可以在不存储 S 矩阵的情况下直接在 SRAM 中施加线性偏置,无需额外开销。
- 长上下文 N > 1M:结合 Ring Attention 和 FlashAttention 的 SP (sequence parallel) 实现 N=1M+ 训练
- 异构硬件适配:FlashAttention 已被移植到 AMD ROCm (AITer)、Intel GPU、Apple Silicon
- 编译器集成:MLIR/Triton 编译器自动生成 IO-aware attention kernel,减少手工 CUDA 开发成本
- 专用硬件加速:Google TPU v5p、Cerebras CS-3 针对 FlashAttention 的数据流专门优化
- Dao, T., et al. "FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness." NeurIPS 2022.
- Dao, T. "FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning." ICML 2023.
- Shah, J., et al. "FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision." 2024.
- Rennie, S.J., et al. "Efficient Block Algorithms for Neural Sequence Labeling." 2014.
- Milakov, M., & Gimelshein, N. "Online normalizer calculation for softmax." 2018.
- NVIDIA. "Hopper Architecture In-Depth." CUDA Toolkit Documentation 12.x.
以 GPT-3 (d=12288, N=2048) 为例,S 矩阵占用 2048×2048×4B = 16MB。每个 Attention Head 需要把 S 和 P 总共 32MB 写入 HBM,再读回用于下一步计算。当 batch size 增大、head 数增多时,IO 带宽成为瓶颈。
GPU 的计算吞吐与带宽之间存在巨大鸿沟。以 A100 为例:
这意味着每加载一个 32-bit 浮点数,GPU 可以执行 156 次浮点运算。但标准 Attention 每 1 次乘加操作需要多次 HBM 访问(加载 Q/K、写 S、读 S、读 P...)。Attention 是 IO-bound 算子,而非 compute-bound。
FlashAttention 的核心洞见:既然 Attention 是 IO-bound 的,优化目标就应该从 FLOPs 最小化 转为 HBM 访问次数最小化。
二、数学基础:Online Softmax 与数值稳定性
FlashAttention 的关键数学突破在于:无需一次性计算完整 N×N 矩阵,就能正确计算 softmax 结果。
2.1 经典 Softmax 的两轮遍历
给定向量 $x \in \mathbb{R}^N$,softmax 定义为:
\[\text{softmax}(x)_i = \frac{e^{x_i}}{\sum_{j=1}^{N} e^{x_j}}$\]
为数值稳定,需减去最大值:
\[\text{softmax}(x)_i = \frac{e^{x_i - m}}{\sum_{j=1}^{N} e^{x_j - m}}$\]
其中 $m = \max_{j}(x_j)$。经典实现需要两轮遍历:第 1 轮找最大值 $m$,第 2 轮计算归一化分母和输出。
2.2 Online Softmax:增量合并算法
FlashAttention 的核心是 Online Softmax(Rennie et al., 2014;Milakov & Gimelshein, 2018),允许增量式计算 softmax。
假设已将序列分为两个子块 $[x^{(1)}, x^{(2)}]$,分别计算后如何合并?
已知:
全局最大值 $m = \max(m^{(1)}, m^{(2)})$,则全局分母:
\[\ell = \ell^{(1)} \cdot e^{m^{(1)}-m} + \ell^{(2)} \cdot e^{m^{(2)}-m}$\]
全局输出中,子块 1 部分的系数变为 $\frac{e^{x_i^{(1)}-m}}{\ell} = \frac{e^{x_i^{(1)}-m^{(1)}}}{\ell^{(1)}} \cdot \frac{\ell^{(1)} \cdot e^{m^{(1)}-m}}{\ell}$
即:旧的输出值需要乘以一个修正因子 $e^{m^{(1)}-m} \cdot \ell^{(1)}/\ell$。
这就是 FlashAttention 中 "re-scaling" 的核心数学原理。
三、FlashAttention 算法详解:Tiling + Recomputation
3.1 算法伪代码
FlashAttention(Q, K, V):
// Q,K,V shape: (N, d) in HBM
O = zeros(N, d) in HBM // 输出
l = zeros(N) in HBM // softmax 归一化项
m = -inf(N) in HBM // 行最大值
// 分块大小(由 SRAM 容量决定)
Br = min(d, max SRAM/(4*d)) // Q 的块
Bc = ... // K,V 的块
// 将 Q, K, V 按行/列分块
Q_i = split(Q, Br) // T_r = ceil(N/Br) 块
K_j, V_j = split(K, Bc), split(V, Bc) // T_c = ceil(N/Bc) 块
for j in range(T_c):
load K_j, V_j from HBM to SRAM
for i in range(T_r):
load Q_i, O_i, l_i, m_i from HBM to SRAM
// 计算 S_ij = Q_i * K_j^T (Br × Bc)
S_ij = Q_i @ K_j^T
// Online softmax 更新
m_ij_new = rowmax(S_ij)
m_i_new = max(m_i, m_ij_new)
P_ij = exp(S_ij - m_i_new) // 局部 softmax
l_i_new = exp(m_i - m_i_new) * l_i + rowsum(P_ij)
// 缩放旧输出,加上新贡献
O_i = (l_i * exp(m_i - m_i_new) / l_i_new) * O_i
+ (1/l_i_new) * P_ij @ V_j
// 记录更新后的统计量
l_i = l_i_new
m_i = m_i_new
write O_i, l_i, m_i back to HBM
return O
3.2 关键优化:无需保存 N×N Attention 矩阵
标准实现需要将 S (N×N) 写入 HBM,占用 O(N²) 空间。FlashAttention 在 SRAM 中直接计算 softmax 和 V 的乘积,只将 O (N×d)、归一化项 l 和行最大值 m 写回 HBM。
HBM 访问复杂度对比:
3.3 Softmax Rescaling 的正确性
第 3 步中出现的缩放因子 $\exp(m_i - m_i_{new})$ 至关重要。代码如下:
// 当发现新块的最大值更大时,需要修正之前的结果
float old_m = m[i]; // 旧行最大值
float new_m = fmaxf(old_m, block_max); // 新行最大值
float scale_old = expf(old_m - new_m); // 缩放旧值
float scale_new = expf(block_max - new_m); // 缩放新贡献
// 修正已计算的输出
O[i] = O[i] * (l[i] * scale_old / l_new) + (scale_new / l_new) * P_block * V_block;
l[i] = l[i] * scale_old + l_block * scale_new;
m[i] = new_m;
这保证了在不存储完整注意力矩阵的前提下,最终输出与标准 Attention 在数学上完全等价——不是近似,是 bit-exact 相同的结果。
四、FlashAttention-2:极致利用 GPU 并行性
2023 年 Dao 发布的 FlashAttention-2(FL2)通过三个方向进一步压榨 Ampere 架构性能:
4.1 减少非乘加操作 (non-matmul FLOPs)
标准算法中,每计算一个分块都需要执行 softmax 修正(指数、乘法占比高)。FL2 将 softmax 操作从 O(N²) 的频率降低到 O(Nd) 的级别:
// FL2:延迟 rescaling 到所有 K,V 块计算完成
for j in range(T_c):
for i in range(T_r):
// 只计算 P_ij 和局部统计量,累积但不立即修正 O_i
accumulate(O_raw_i, l_i, m_i, K_j, V_j)
// 一轮 K,V 全部完成后,做一次 rescaling
O_i = O_i / l_i
这将非乘加操作减少约 5-8×。
4.2 Warp 间并行优化
FL2 改变了线程块内的并行策略:
# 并行分配策略
# block_handle: 每个 block 处理 Br 个 Q 行
# warp_handle: 同一 block 内 warp_i 处理行 [i*Br/wp : (i+1)*Br/wp]
# wp = warps per block
这种策略减少了 warp 间的同步开销,同时增加了数据复用效率。
4.3 循环重排 (Loop Reordering)
FL2 将外层循环从 "K 外 / Q 内" 改为 "Q 外 / K 内",使得 Q 块只需载入一次 SRAM,K/V 块在流动刷新:
# 原始:for K_blocks → for Q_blocks (每加载 K 就循环所有 Q)
# FL2: for Q_blocks → for K_blocks (每加载 Q 就遍历所有 K)
为了在 Q 外层循环中保持 Online Softmax 的正确性,FL2 修改了累积策略,将 softmax 合并从在线式改为批处理式。
五、FlashAttention-3:Hopper 架构原生优化
2024 年的 FlashAttention-3(FL3)专为 NVIDIA Hopper GPU (H100) 设计,利用三项硬件特性:
5.1 Warpgroup Level MMA ( wgmma )
H100 引入的 wgmma 指令允许直接在共享内存(SRAM)和 Tensor Core 之间传输数据,eliminate 寄存器中转步骤。FL3 通过 wgmma 流水线化计算:
// wgmma 流水线阶段
stage 1: wgmma_async SRAM → Tensor Core (K, V)
stage 2: wgmma_async SRAM → Tensor Core (Q)
stage 3: wgmma_commit_sync 等待计算完成
stage 4: prefetch下一块数据
5.2 FP8 低精度计算
H100 的 FP8 Tensor Core 提供 3,958 TFLOPS(对比 FP16 的 989 TFLOPS)。FL3 在 Attention 计算中使用 FP8 精度存储 QK^T 和 PV,但由于 softmax 对精度敏感,仍使用 FP32 计算统计量(m, l):
// FL3 混合精度策略
__nv_fp8_e4m3 QK_T; // 低精度存储注意力分数
float m, l; // FP32 用于 softmax 统计量(数值稳定性)
half O; // FP16/BF16 输出精度可选
5.3 异步拷贝与 TMA (Tensor Memory Accelerator)
H100 的 TMA 引擎可直接从 HBM 传输数据到 SRAM,无需 SM 主动拷贝:
// TMA 异步预取下一块 K, V
void cp_async_bulk(void *smem, const void *gmem, size_t size);
cp_async_bulk(sram_buffer, &K[next_block], block_size);
cp_async_bulk_commit_group();
cp_async_bulk_wait_group<0>(); // 等待预取完成
TMA 与 wgmma 配合,实现从 HBM → SRAM → Tensor Core 的全流水线化,最大化带宽利用率。
六、工程实践:CUDA 内核实现模式
6.1 Block 大小选择
SRAM 容量决定分块大小。A100 有 192KB SMEM 每 SM,FL2 的典型配置:
# FlashAttention-2 Triton 内核配置示例
@triton.jit
def _fwd_kernel(
Q, K, V, Out,
stride_qz, stride_qh, stride_qm, stride_qk,
BLOCK_M: tl.constexpr, # 通常为 128 (Q 块)
BLOCK_N: tl.constexpr, # 通常为 128 (K,V 块)
BLOCK_DMODEL: tl.constexpr,
...
):
# BLOCK_M * BLOCK_DMODEL + 2 * BLOCK_N * BLOCK_DMODEL ≤ SMEM_SIZE
# 128 × 64 + 2 × 128 × 64 = 24,576 bytes (FP16)
# 剩余空间:softmax 中间结果、m、l 向量等
序列长度 N 很大时,128×128 的块能充分利用 SMEM;N 较小时,Bc 可增加到 250。
6.2 自动调优器 (Auto-Tuner)
实际部署中,最优的 BLOCK_M、BLOCK_N、num_warps 是通过自动调优选择的:
# flash_attn/flash_fwd_launch.h 简化示意
template<int Headdim>
void run_mha_fwd_(Flash_fwd_params ¶ms, cudaStream_t stream) {
if (params.is_sm80 || params.is_sm90) {
// Ampere/Hopper: 使用 wgmma
dim3 grid(params.b * params.h, params.seqlen_q / BLOCK_M);
flash_fwd_hoppersgp_kernel<<<grid, ...>>>(params);
} else {
// Turing: 使用 legacy CUDA
flash_fwd_additivemasks_kernel<<<grid, ...>>>(params);
}
}
6.3 内存占用对比表
| 算法 | HBM 辅助空间 | SRAM 峰值 | 支持序列长度 |
|---|---|---|---|
| 标准 Attention | O(N²) | O(Bc·d) | ~2K (A100 80GB) |
| FlashAttention | O(N) | O(Bc·d) | ~65K |
| FlashAttention + 双向分块 | O(N) | O(Bc·d) | ~1M |
| FlashAttention-3 FP8 | O(N) | O(Bc·d/2) | ~130K |
七、性能基准与对比分析
7.1 训练吞吐量 (A100, GPT-3 175B)
在 Dao 等作者的基准测试中,Standard PyTorch Attention vs FlashAttention:
| 序列长度 | 标准 Attention | FlashAttention-2 | 加速比 |
|---|---|---|---|
| 1,024 | 100 TFLOPS | 195 TFLOPS | 1.95× |
| 2,048 | 135 TFLOPS | 275 TFLOPS | 2.04× |
| 4,096 | 155 TFLOPS | 310 TFLOPS | 2.00× |
| 8,192 | 170 TFLOPS | 330 TFLOPS | 1.94× |
| 16,384 | 175 TFLOPS | 340 TFLOPS | 1.94× |
7.2 端到端训练加速 (GPT-3)
| 训练配置 | Standard | FlashAttention | 时间 |
|---|---|---|---|
| GPT-3 7B | 基准 | 1.5× | 40% 时间 |
| GPT-3 13B | 基准 | 1.8× | 44% 时间 |
| GPT-3 175B | 基准 | 2.1× | 52% 时间 |
7.3 H100 FlashAttention-3 性能
H100 上 FL3 使用 FP8:
| 序列长度 | FA-2 FP16 | FA-3 FP8 | 加速比 |
|---|---|---|---|
| 4,096 | 330 TFLOPS | 660 TFLOPS | 2.0× |
| 8,192 | 340 TFLOPS | 700 TFLOPS | 2.06× |
| 16,384 | 340 TFLOPS | 720 TFLOPS | 2.12× |
八、应用场景与工程权衡
8.1 何时使用 FlashAttention
适合场景:
不适合场景:
8.2 工程注意事项
九、总结与展望
FlashAttention 的核心贡献不是提出新的算子,而是重新定义了 Attention 优化的目标函数:
标准视角:Attention 是 compute-bound → 优化 FLOPs
FlashAttention 视角:Attention 是 IO-bound → 优化 HBM 访问
这种从硬件感知出发的设计哲学正在重塑整个 GPU 算子生态。从 RMSNorm 的融合、SwiGLU 的 TMA 优化,到 VLLM 的 PagedAttention(本质上也是 IO-aware 的管理策略),都遵循着同一范式。
展望方向:
在软件定义硬件的时代,理解 IO 墙、优化数据局部性、将算法与架构协同设计——这正是 FlashAttention 留给工程师的最大启示。

发表评论 取消回复