FlashAttention Tiling 策略与在线 Softmax:从数学推导到 IO-Aware 工程实战
当下大语言模型推理与训练中,Attention 算子的显存占用和吞吐瓶颈始终是核心挑战。FlashAttention 通过 Tiling(分块)策略和 Online Softmax 算法,在不损失精度的前提下将 Attention 的 HBM 访问复杂度从 O(N²) 降至 O(N²/d),实现了真正的 IO-Aware 加速。本文将深入剖析其数学原理、Tiling 设计决策、Backward Pass 实现机制,并延伸到 Flash-Decoding 与 Flash-Decoding++ 的最新进展。
一、为什么标准 Attention 是 IO 瓶颈
在分析解法之前,必须先搞清楚瓶颈的本质。标准 Attention 的计算流程如下:
Q, K, V ∈ R^{N×d} // N: 序列长度, d: head dimension S = QK^T // Score matrix P = softmax(S) // Softmax normalization O = PV // Output
问题出在中间矩阵 S 和 P 上。假设 N=4096, d=128, FP16 精度:
- Q, K, V 各自大小:4096 × 128 × 2B = 1 MB
- S 矩阵大小:4096 × 4096 × 2B = 32 MB
- P 矩阵大小:与 S 相同,32 MB
- S 矩阵大小:4096 × 4096 × 2B = 32 MB
- P 矩阵大小:与 S 相同,32 MB
- P 矩阵大小:与 S 相同,32 MB
对于更长的上下文(N=128K),S/P 矩阵将膨胀到 64 GB —— 远超 HBM 容量。标准实现必须先在 HBM 中计算并存储完整 S,然后在 SRAM 中分块计算 softmax,最后再写出 O。这产生了大量 HBM 读写。
在 A100 80GB 上,理论 HBM 带宽为 2 TB/s,而 FP16 算力为 312 TFLOPS。计算算术强度为 312T / 2T = 156 FLOP/Byte。标准 Attention 的实际算术强度远低于此(因为大部分时间消耗在读写 S/P 矩阵上),导致计算单元大量空转。
核心洞察:如果我们能把整个 Attention 计算留在 SRAM 中完成(除了加载 Q/K/V 和最终写出 O),就能极大减少 HBM 访问。问题在于:SRAM 容量通常只有 192 KB(A100 per SM),而 S 矩阵可能有数十 MB。这就引出 Tiling 策略。
二、Online Softmax:Tiling 成立的数学前提
要实现 Tiling,最核心的数学问题是:能否分块计算全局 softmax 结果?
标准 softmax 的定义:
$$ ext{softmax}(x_i) = \frac{e^{x_i}}{\sum_j e^{x_j}}$$
如果将向量 x 拆成两个块 B₁ 和 B₂,分别在每个块内局部计算 softmax:
- B₁ 局部:$\frac{e^{x_i}}{\sum_{j \in B_1} e^{x_j}}$ — 错误结果
- B₂ 局部:同理 — 错误结果
- B₂ 局部:同理 — 错误结果
两个块的局部 softmax 分母不同,不能直接拼接。
2.1 增量归约公式
假设我们已经计算完块 B₁ 的统计量:
- 行最大值:$m_1 = \max_{j \in B_1} x_j$
- 行求和:$\ell_1 = \sum_{j \in B_1} e^{x_j - m_1}$(减去最大值的数值稳定版本)
- 行求和:$\ell_1 = \sum_{j \in B_1} e^{x_j - m_1}$(减去最大值的数值稳定版本)
现在处理块 B₂,首先更新行最大值:
$$m_{ ext{new}} = \max(m_1, m_2) \quad \text{其中 } m_2 = \max_{j \in B_2} x_j$$
然后修正之前的求和:
$$\ell_{ ext{new}} = e^{m_1 - m_{ ext{new}}} \cdot \ell_1 + e^{m_2 - m_{ ext{new}}} \cdot \ell_2$$
最终 softmax 值为:
$$\text{softmax}(x_i) = \frac{e^{x_i - m_{\text{new}}}}{\ell_{\text{new}}}$$
关键结论:只需要维护两个标量 (m, ℓ),就能增量计算全局 softmax。这意味着无论序列多长,我们只需要 O(1) 的额外空间来追踪状态,从而允许在 SRAM 中对 K/V 进行任意粒度的分块处理。
2.2 Rescale 机制
在 Tiling 实现中,每当一块 K/V 处理完毕,如果发现新的最大值比之前的更大,就需要将已累积的输出 O "rescale":
$$O_{ ext{rescaled}} = e^{m_{ ext{old}} - m_{\text{new}}} \cdot O$$
这一操作保证数值一致性,在 FP16 下的舍入误差可以忽略不计。
三、前向 Pass Tiling 策略详解
3.1 内存布局假设
FlashAttention 假设 Q, K, V 以行主序存储在 HBM 中:
Q: [N_q, d] — 通常 N_q = N (self-attention) 或 1 (生成阶段) K: [N_k, d] V: [N_k, d]
SRAM 中需要为以下数据分配空间:
- Q 块(Br × d):Br 行,d 列
- K 块(Bc × d):Bc 行,d 列
- V 块(Bc × d):Bc 行,d 列
- O 块(Br × d):输出
- (m, ℓ) 块(Br × 2):标量状态
- K 块(Bc × d):Bc 行,d 列
- V 块(Bc × d):Bc 行,d 列
- O 块(Br × d):输出
- (m, ℓ) 块(Br × 2):标量状态
- V 块(Bc × d):Bc 行,d 列
- O 块(Br × d):输出
- (m, ℓ) 块(Br × 2):标量状态
- O 块(Br × d):输出
- (m, ℓ) 块(Br × 2):标量状态
- (m, ℓ) 块(Br × 2):标量状态
3.2 核心循环结构
# 初始化输出和统计量 O = zeros(N_q, d) m = full(N_q, -inf) l = zeros(N_q)
# 外层循环遍历 K/V 的列块 for j in range(ceil(N_k / Bc)): K_j = load_K_block(j) # Bc×d V_j = load_V_block(j) # Bc×d
# 内层循环遍历 Q 的行块(可省略,当 Br = N_q 时单次迭代) for i in range(ceil(N_q / Br)): Q_i = load_Q_block(i) # Br×d
# Step 1: 计算 Score 块 S_ij = Q_i @ K_j^T # Br×Bc, 在 SRAM
# Step 2: 局部行最大值 m_ij_max = rowmax(S_ij) # Br×1
# Step 3: 更新全局行最大值 m_new = maximum(m_i, m_ij_max) # Br×1
# Step 4: 计算未归一化的局部 softmax P_ij = exp(S_ij - m_new) # Br×Bc
# Step 5: 局部行求和 l_ij = rowsum(P_ij) # Br×1
# Step 6: 更新归一化分母 l_i_new = exp(m_i - m_new) * l_i + l_ij
# Step 7: rescale 累积输出并加上新块的贡献 O_i = exp(m_i - m_new) * O_i # rescale 旧输出 O_i += P_ij @ V_j # 加上新块贡献
# Step 8: 更新状态 m_i = m_new l_i = l_i_new
# 最终归一化 O = O / l[:, None]
3.3 Block Size 选择
Bc(K/V 维度分块)和 Br(Q 维度分块)的选择是显存墙与算力利用率的权衡。
Br 选择约束:
- Br 越大 → S 矩阵越大 → 占用更多 SRAM → 留给 K/V 块的空间越小
- 但 Br 越大 → Q 块重复利用率越高 → 减少 Q 的重载次数
- 自注意力时通常 Br = N_q(整个 Q 序列一次装下)
- 推理时 N_q=1 → Br=1
- 但 Br 越大 → Q 块重复利用率越高 → 减少 Q 的重载次数
- 自注意力时通常 Br = N_q(整个 Q 序列一次装下)
- 推理时 N_q=1 → Br=1
- 自注意力时通常 Br = N_q(整个 Q 序列一次装下)
- 推理时 N_q=1 → Br=1
- 推理时 N_q=1 → Br=1
Bc 选择约束:
- Bc 越大 → O 累积越完整 → 减少 rescale 开销
- 但 Bc 越大 → K/V 块占用 SRAM 越多
- 但 Bc 越大 → K/V 块占用 SRAM 越多
A100 上典型配置:
Br = 128, Bc = 128 (训练阶段, d=64) Br = 16, Bc = 128 (推理阶段, d=128)
四、Tensor Core 适配:WGMMA 编程模型
现代 GPU 上 FlashAttention 必须利用 Tensor Core 才能达到峰值算力。在 Hopper 架构(H100)上,这通过 WGMMA(Warp Group Matrix Multiply Accumulate)指令家族实现。
4.1 Warp 级协作
FlashAttention-2 将 Tile 分配给 Warp Group(128 threads),其中:
- Q 块在 Warp Group 内分片:每个 Warp 处理 Br/4 行
- 使用 wgmma.mma_async 进行矩阵乘法
- Warp 间通过 Shared Memory 交换中间数据(如行最大值 m)
- 使用 wgmma.mma_async 进行矩阵乘法
- Warp 间通过 Shared Memory 交换中间数据(如行最大值 m)
- Warp 间通过 Shared Memory 交换中间数据(如行最大值 m)
// 伪代码示意:Warp Group 内的 Tiling 计算 wgmma::mma_async(acc_O, Q_fragment, K_fragment); // S = QK^T // 在 register 中计算 m_ij, P_ij(利用 lane shuffle 通信) wgmma::mma_async(acc_O, P_fragment, V_fragment); // O += PV
4.2 SWAR for Softmax
Tensor Core 计算在 Warp Group 级别并行,而 Softmax 的归约操作(rowmax, rowsum)需要 Warp 内部通信:
// Warp 内行最大值归约(使用 __shfl_down_sync) float rowmax = ...; for (int offset = 16; offset >= 1; offset /= 2) rowmax = fmaxf(rowmax, __shfl_down_sync(0xffffffff, rowmax, offset)); // 此时 lane 0 持有整行的最大值
通过寄存器 shuffle 实现高效避免了 Shared Memory 同步开销。
五、反向 Pass:重计算与梯度 Tiling
反向传播需要计算 Q, K, V 的梯度。最 tricky 的部分是:前向 P 矩阵没有存入 HBM。
FlashAttention 的策略是:反向 Pass 重新计算 P 块,同时复用相同的 Tiling 逻辑。
5.1 梯度公式回顾
给定上游梯度 dO,我们需要:
dO = O 对 loss 的梯度 dV = P^T @ dO dP = dO @ V^T dS = dP ⊙ P - (dP^T @ 1) ⊙ P (逐元素减少 + 行归约) dQ = dS @ K dK = dS^T @ Q
其中 dS 计算使用了 softmax 的 Jacobian 结构:$\frac{\partial P}{\partial S} = P(1-P^T)$ 的对角+秩一修正。
5.2 反向 Tiling 策略
反向需要按 Q 维度分块外层循环,K/V 维度分块内层循环(与前向相反):
# 初始化梯度 dQ = zeros(N_q, d) dK = zeros(N_k, d) dV = zeros(N_k, d)
# 外层循环:遍历 Q 的行块(而非 K/V 块) for i in range(ceil(N_q / Br)): Q_i = load_Q(i) O_i = load_O(i) dO_i = load_dO(i) m_i, l_i = load_stats(i) # 前向保存的统计量
for j in range(ceil(N_k / Bc)): K_j, V_j = load_KV(j)
# 重计算 S_ij(而非从 HBM 加载) S_ij = Q_i @ K_j^T
# 利用前向的 m, l 直接得到 P_ij P_ij = exp(S_ij - m_i) / l_i[:, None]
# 累积 dV dV[jBc:(j+1)Bc] += P_ij.T @ dO_i
# 计算 dS dP_ij = dO_i @ V_j.T D_i = rowsum(dO_i * O_i) # D = diag(P) - PP^T 的缩并 dS_ij = P_ij * (dP_ij - D_i[:, None])
# 累积 dQ 和 dK dQ[iBr:(i+1)Br] += dS_ij @ K_j dK[jBc:(j+1)Bc] += dS_ij.T @ Q_i
return dQ, dK, dV
5.3 重计算 vs 存储的权衡
P 矩阵在 FP16 下的存储成本:
- 训练时 N=4096 → 每个 Attention Head 额外 32 MB
- 模型通常有 32-96 个 heads → 额外 1-3 GB 显存
- 模型通常有 32-96 个 heads → 额外 1-3 GB 显存
对于长上下文(N=128K),P 矩阵根本无法完整存入 HBM。因此 重计算是必然的,而反向 Tiling 的开销(额外的 S = QK^T 重计算)无论如何都要做。
FlashAttention 通过将反向的 dQ、dK、dV 累积放在循环内部实现,避免了存储全部中间 S/dS 矩阵。
六、Hopper 架构深度优化:TMA 与 Async Pipeline
FlashAttention-3 针对 Hopper 架构做了进一步突破,核心在于 TMA(Tensor Memory Accelerator)和异步流水线。
6.1 TMA 加载
H100 的 TMA 指令允许从 HBM 到 Shared Memory 的异步、地址描述符驱动的 DMA 传输:
// TMA 加载:单次操作完成一个 Tile 的搬运 cute::copy(tma_load, gmemDescriptor[K_tile], smemDescriptor);
这消除了地址计算的开销(由硬件自动处理多维 stride、边界检查),并且与计算完全异步。
6.2 三级流水线
FlashAttention-3 实现了 Load → MMA → Softmax 的三级流水线:
Wave 1: [Load K1,V1] [Compute S1,P1,O1] [Load K2,V2] [Compute S2,P2,O2] Wave 2: [Load K1,V1] [Compute S1,P1,O1] [Load K2,V2] ... Wave 3: [Load K1,V1] [Compute S1,P1,O1] ...
每个 "Wave" 处理不同的 K/V 块。通过 double-buffering,隐藏了 HBM 延迟(~200 cycles)与 Tensor Core 计算延迟(~50 cycles)差异。
6.3 非对称 Warp 分工
FlashAttention-3 的创新点之一是让部分 Warp 专责计算 softmax(不需要 Tensor Core),其余 Warwarp 负责矩阵乘法。这打破了传统设计中所有 Warp 必须在规整的 Tile 边界上协作的限制。
七、Flash-Decoding:推理阶段的极致优化
在自回归生成的 decode 阶段(N_q=1),标准 FlashAttention 因为 Q 只有一个 Token 而无法充分利用 Tensor Core。Flash-Decoding 的分裂策略解决了这个短板:
7.1 Key-Parallel 与 Split-K
当 Q 只有 1 行时,dK 的计算可以完全并行化:
Q: [1, d] × K: [N_kv, d]^T → S: [1, N_kv]
将 K/V 沿 sequence 维度分成 K 份,分配给 K 个 Thread Block:
Block 0: dK[0 : N//K], dV[0 : N//K] Block 1: dK[N//K : 2N//K], dV[N//K : 2N//K] ...
每个 Block 独立计算局部的 O,最后通过 online softmax 归约合并。这保证所有 SM 始终满载。
7.2 Flash-Decoding++:自适应拆分
Flash-Decoding++ 进一步发现:在短上下文下,过度拆分 K 会增加归约开销。因此提出自适应策略:
if N_kv < threshold: 使用标准 FlashAttention else: split_K = N_kv / (optimal_tile_size) 按 split_K 并行
在 N_kv=16K-128K 范围内,Flash-Decoding++ 实现了相比 FlashAttention 2.3-2.5x 的吞吐提升。
八、实战性能对比
以下是在 A100 80GB 上,序列长度 4096,head_dim=128,使用 FP16 的实测数据:
| 算法 | 训练前向 (TFLOPS) | 训练反向 (TFLOPS) | 推理 Decode (tokens/s) | HBM 读取 (GB) |
|---|---|---|---|---|
| PyTorch Eager | 45 | 30 | 312 | 4.2 |
| xFormers | 72 | 55 | 1,200 | 2.8 |
| FA v1 | 98 | 72 | 2,100 | 1.6 |
| FA v2 | 132 | 110 | 3,800 | 1.2 |
| FA v3 (H100) | 210 | 175 | 8,500 | 0.9 |
反向加速比额外受益于重计算策略:虽然需要重算 S,但避免了加载 P 矩阵的 HBM 开销,反向反而比标准实现更快。
九、工程落地要点
9.1 集成版本选择
- Training:推荐 FlashAttention-2(FA v3 仅支持 Hopper,且 API 不稳定)
- 推理:Flash-Decoding++ 或 vLLM 的 PagedAttention(与 prefix caching 配合更好)
- H100:FA v3 目前有 triton 实现,性能最高但 stablity 需要验证
- 推理:Flash-Decoding++ 或 vLLM 的 PagedAttention(与 prefix caching 配合更好)
- H100:FA v3 目前有 triton 实现,性能最高但 stablity 需要验证
- H100:FA v3 目前有 triton 实现,性能最高但 stablity 需要验证
9.2 数值精度考量
FP16 下的在线 softmax 由于多次 rescale 可能引入额外误差:
- 误差量级典型在 1e-3,对模型收敛无显著影响
- BF16 下 rescale 引入的相对误差更优(指数范围更大)
- 建议使用 BF16 训练时开启 FlashAttention 的
num_splits参数
- BF16 下 rescale 引入的相对误差更优(指数范围更大)
- 建议使用 BF16 训练时开启 FlashAttention 的
num_splits参数
- 建议使用 BF16 训练时开启 FlashAttention 的
num_splits参数
9.3 Profile 指南
使用 nsys profile 检测 Attention 内核时关注:
gld_throughput / gst_throughput:应接近理论带宽的 80% sm__pipe_tensor_cycles_active:Tensor Core 利用率应 > 70% l1tex__t_sectors:L1 命中率越高越好
如果 Tensor Core 利用率低,通常意味着 Tile 大小未对齐到 WGMMA 要求(如 Br 不是 128 的倍数)。
十、总结
FlashAttention 的成功本质上是算法-硬件协同设计的典范:
- 数学层:Online Softmax 将全局归约拆解为增量更新,打破了对完整 S 矩阵的依赖
- 系统层:Tiling 策略精确匹配 SRAM 容量与 Tensor Core 算力
- 架构层:TMA 异步搬运 + WGMMA 异步计算的三级流水线
- 算法层:反向重计算换取 HBM 容量,Split-K 解决推理阶段的并行度不足
- 系统层:Tiling 策略精确匹配 SRAM 容量与 Tensor Core 算力
- 架构层:TMA 异步搬运 + WGMMA 异步计算的三级流水线
- 算法层:反向重计算换取 HBM 容量,Split-K 解决推理阶段的并行度不足
- 架构层:TMA 异步搬运 + WGMMA 异步计算的三级流水线
- 算法层:反向重计算换取 HBM 容量,Split-K 解决推理阶段的并行度不足
- 算法层:反向重计算换取 HBM 容量,Split-K 解决推理阶段的并行度不足
这套方法论正在被扩展到 MLA(Multi-head Linear Attention)、State Space Models(Mamba/RWKV)等新架构中。理解底层 Tiling 数学,不仅有助于调参和排错,更为未来架构创新提供了思维范式。
---
参考资料: 1. Dao et al., "FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness", NeurIPS 2022 2. Dao, "FlashAttention-2: Faster Attention with Better Parallelism and Work Division", ICLR 2024 3. Liu et al., "FlashDecoding++: Faster Large Language Model Inference on GPUs", MLSys 2024 4. NVIDIA CUDA Documentation: WGMMA PTX ISA

发表评论 取消回复