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 Eager45303124.2
xFormers72551,2002.8
FA v198722,1001.6
FA v21321103,8001.2
FA v3 (H100)2101758,5000.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

点赞(0) 打赏

评论列表 共有 0 条评论

暂无评论
立即
投稿

微信公众账号

微信扫一扫加关注

发表
评论
返回
顶部