从第一性原理推导 Flash Attention:IO 复杂度、Tiling 策略与在线 Softmax 的数学工程
注意力机制是 Transformer 架构的核心计算瓶颈。当序列长度达到 32K、128K 甚至更长时,标准 Attention 的 O(N²) 内存复杂度和 HBM 带宽限制成为推理服务的根本瓶颈。Flash Attention 不是一篇"调优论文",而是一次从存储层次结构出发的算法重新设计。本文将从第一性原理出发,推导 Flash Attention 的核心数学变换,分析其 IO 复杂度下界,并讨论工程落地中的关键细节。
一、标准 Attention 的瓶颈在哪里
标准 Scaled Dot-Product Attention 的计算公式:
Attention(Q, K, V) = softmax(QK^T / √d) · V
其中 Q, K, V ∈ ℝ^(N×d),N 是序列长度,d 是 head 维度。
问题出在中间矩阵 S = QK^T ∈ ℝ^(N×N)。当 N=128K,d=128 时:
- S 矩阵大小:128000² × 2 bytes (fp16) = 31.25 GB
- P 矩阵(softmax 后)同样大小
- HBM 到 SRAM 的数据搬运成为瓶颈
以 A100 为例:HBM 带宽 2 TB/s,SRAM 带宽 19 TB/s(每 SM 约 19TB/s)。算法算术强度(Arithmetic Intensity)决定了计算是 compute-bound 还是 memory-bound。
二、IO 复杂度下界的理论推导
2.1 分块计算的必要性
传统实现需要把整个 S 矩阵写入 HBM 再做 softmax,然后再读回来与 V 相乘。Flash Attention 的核心洞察是:我们不需要在 HBM 中保存完整的 S 矩阵。
将 Q, K, V 沿序列维度分块:
Q = [Q_1, Q_2, ..., Q_T] Q_i ∈ ℝ^(B_r × d)
K = [K_1, K_2, ..., K_T] K_j ∈ ℝ^(B_c × d)
V = [V_1, V_2, ..., V_T] V_j ∈ ℝ^(B_c × d)
其中 T = ⌈N/B_r⌉ = ⌈N/B_c⌉。
2.2 IO 复杂度推导
对于每个 Q_i 块,我们需要遍历所有 K_j, V_j 块来计算输出 O_i。传统算法的 IO 复杂度:
- 标准 Attention:O(N²d + N²) = O(N² · max(Nd, N)/M),实际简化为 O(N² · d / M) 次 HBM 访问
- Flash Attention:O(N² · d² · M / ...),经过优化的复杂度为 O(N²d²M⁻¹)
其中每个核心操作: 1. 加载 K_j, V_j 到 SRAM(B_c × d × 2 个元素) 2. 计算 S_ij = Q_i K_j^T(B_r × B_c) 3. 在线更新 softmax 统计量 4. 累加输出块
Flash Attention 的总 IO 次数为:
IO_std = Θ(N²d)
IO_flash = Θ(N²d²M⁻¹) (M 为 SRAM 大小)
当 M=168KB(A100 每 SM),d=128,d²M⁻¹ = 128²/(168×1024) ≈ 0.095,实现了约 10 倍 IO 减少。
三、在线 Softmax:核心数学变换
这是 Flash Attention 最精妙的部分。softmax 需要知道所有元素的最大值才能归一化,但分块处理时我们已经处理过的块没有"未来"块的信息。
3.1 Softmax 的分块计算公式
对于向量 x = [x_1, x_2, ..., x_T],标准 softmax:
softmax(x)_i = exp(x_i - max(x)) / Σ_j exp(x_j - max(x))
分块处理时,假设已处理到第 j 块,我们有:
- m_j = max(m_{j-1}, max(x_j)) —— 全局最大值
- l_j = exp(m_{j-1} - m_j) · l_{j-1} + Σ_i exp(x_j,i - m_j) —— 修正后的求和
- O_j = exp(m_{j-1} - m_j) · O_{j-1} + exp(x_j - m_j) · V_j —— 修正后的输出
3.2 数学正确性证明
合并引理(Merge Lemma): 给定两个分块统计量 (m_a, l_a, O_a) 和 (m_b, l_b, O_b),合并后的结果为:
m_new = max(m_a, m_b)
l_new = exp(m_a - m_new) · l_a + exp(m_b - m_new) · l_b
O_new = exp(m_a - m_new) · O_a + exp(m_b - m_new) · O_b
这个引理的正确性来自 softmax 的"平移不变性":
softmax(x) = softmax(x - c) 对任意常数 c
当我们发现新的最大值 m_new > m_a 时,之前计算的 O_a 需要乘以 exp(m_a - m_new) 来"修正"。这就是 Flash Attention 能在线(online)完成计算而无需两遍扫描的理论基础。
3.3 数值安全性的实现细节
# 伪代码:在线 softmax 核心
def online_softmax_update(m_prev, l_prev, O_prev, q_block, k_block, v_block):
# 计算当前块的 S 矩阵
S = q_block @ k_block.T / sqrt(d) # (B_r, B_c)
# 当前块行最大值
m_curr = S.max(dim=1, keepdim=True) # (B_r, 1)
# 更新全局最大值
m_new = maximum(m_prev, m_curr)
# 计算 exp(S - m_new),需要两个修正因子
P = exp(S - m_new) # (B_r, B_c)
correction_old = exp(m_prev - m_new) # 修正之前的输出
correction_curr = exp(m_curr - m_new) # 修正当前块
# 更新求和统计量
l_new = correction_old * l_prev + correction_curr * P.sum(dim=1, keepdim=True)
# 更新输出:修正之前输出 + 当前块贡献
O_new = correction_old * O_prev + correction_curr * (P @ v_block)
return m_new, l_new, O_new
关键点:使用 exp(m_prev - m_new) 而非直接计算 exp(m_prev) / exp(m_new),避免数值溢出。这在 m_prev >> m_new 时至关重要。
四、Tiling 策略与 SRAM 分区
4.1 SRAM 空间分配
A100 每 SM 有 192KB 可配置 SRAM(与 L1 cache 共享)。Flash Attention 需要以下缓冲区:
| 缓冲区 | 维度 | 大小(fp16) |
|---|---|---|
| Q_i | B_r × d | B_r × 128 × 2B |
| K_j | B_c × d | B_c × 128 × 2B |
| V_j | B_c × d | B_c × 128 × 2B |
| S_ij | B_r × B_c | B_r × B_c × 2B |
| O_i (部分) | B_r × d | B_r × 128 × 2B |
| 中间统计量 | B_r × K | 可忽略 |
总大小 = (B_r + 2B_c) × d × 2 + B_r × B_c × 2 bytes
在 A100(192KB SRAM)上求解: - d=128, 取 B_r=128, B_c=128 → 总 ≈ 128²×2 + 3×128×128×2 = 98,304 B ≈ 96KB - 实际 Flash Attention 2 使用 B_r=128, B_c=128(d≤128)或 B_r=64, B_c=64(d=256)
4.2 Block 大小的选择逻辑
B_r 和 B_c 的权衡:
- B_r 更大:每次 K/V 加载服务更多 Q 的行,提高算术强度;但需要更多 SRAM 给 Q_i 和 O_i
- B_c 更大:每次加载更多列,减少外层循环次数;但占用更多 SRAM 给 K_j, V_j 和 S_ij
- 最优平衡:使 SRAM 刚好用满,同时保持 B_r ≈ B_c 以平衡内外层循环
4.3 Warp 级并行
Flash Attention 的 Tiling 不仅仅在 block 级别。在 block 内部:
- 将 Q_i 的 B_r 行分配到多个 warp
- 每个 warp 负责若干行的 softmax 计算
- Warp 之间通过 warp shuffle 指令交换 m 和 l 的归约结果
- 使用 Tensor Core 加速矩阵乘法(WGMMA 指令)
Block (128 threads = 4 warps)
├── Warp 0: Q rows [0, 7]
├── Warp 1: Q rows [8, 15]
├── Warp 2: Q rows [16, 23]
└── Warp 3: Q rows [24, 31]
五、反向传播工程
Flash Attention 反向传播同样需要分块策略,且比前向更复杂,因为涉及三个中间梯度:
# 反向传播:已知 dO,求 dQ, dK, DV
# S = QK^T, P = softmax(S), O = PV
D = sum(dO * O, dim=-1) # (N, 1) 缩放因子
for j in range(T): # 外层循环:遍历 K, V 块
K_j, V_j = load_kv(j)
S_ij = Q_i @ K_j.T
P_ij = softmax_online(S_ij) # 仅计算不存储
dV_j += P_ij.T @ dO_i
dP_ij = dO_i @ V_j.T
dS_ij = P_ij * (dP_ij - D_i) # 关键:softmax 的雅可比简化
dQ_i += dS_ij @ K_j
dK_j += dS_ij.T @ Q_i
关键工程决策:
- 不存储 P 矩阵:前向时维护在线统计量,反向时重新计算 S_ij 和 softmax。牺牲计算换取内存,净收益为正(因为瓶颈在 IO 不在计算)
- dQ 需要完整的前向遍历:每个 dQ_i 块需要所有 K_j 的信息,因此 dQ 的计算外层循环在 j,内层在 i
- dK, dV 可以与 dQ 合并计算:通过精心设计的 loop tiling 实现单次前向遍历
六、Flash Attention 2 & 3 的关键改进
6.1 Flash Attention 2:减少非矩阵乘法运算
FA1 的改进方向是 IO 复杂度,FA2 进一步优化计算效率:
- 减少非 FMA 操作:softmax 中的 exp 和除法在 CUDA 中开销大,FA2 通过更优的 warp 分工减少这类操作
- 降低 block 间的同步开销:FA1 中 Q 的外层循环导致每处理一个新 Q block 需要重新加载所有 KV;FA2 将循环改为 K/V 外层 + Q 内层,减少同步
- 更好的 warp 分配:将 softmax 除法和矩阵乘分配给不同 warp 执行,实现流水线并行
# FA2 伪代码结构:K, V 外层循环,Q 内层循环
for j in range(num_kv_blocks):
K_j, V_j = load(j) # 加载到 SRAM
for i in range(num_q_blocks): # 内层 Q 循环
Q_i = load(i)
S = Q_i @ K_j.T
m, l, O_i = online_softmax_update(m, l, Q_i, K_j, V_j)
store(O_i)
6.2 Flash Attention 3:利用 Hopper 架构特性
Hopper(H100)引入了 TMA(Tensor Memory Accelerator)和 WGMMA 指令:
- TMA 异步拷贝:硬件加速的 HBM↔SRAM DMA,无需显式 shared memory 管理
- WGMMA 指令:直接在 shared memory 上执行矩阵乘法,减少寄存器↔SRAM 的搬运
- FP8 支持:利用 Hopper 的 FP8 Tensor Core,带宽翻倍
七、工程实践中的关键经验
7.1 与 Paged Attention 的协同
在实际推理引擎中(如 vLLM),Flash Attention 与 Paged Attention 需配合工作:
- Paged Attention 管理 KV Cache 的物理页分配
- Flash Attention 负责单个 request 内多个 chunk 的 attention 计算
- 挑战:不规则的页偏移需要额外的索引查找,影响 Flash Attention 的合并访存
7.2 Causal Mask 的 Causal Tiling
解码阶段的 Causal Mask 可以优化 tiling 策略:
for j in range(i + 1): # 只有 j ≤ i 的块有效
# 处理 K_j, V_j
不需要计算 j > i 的块,可以实现约 2 的计算节省(实际约 1.5-1.8,因为还在加载中做判断)。
7.3 性能调优参数
| 参数 | 典型值 | 影响 |
|---|---|---|
| BLOCK_M | 64-128 | Q block 行数 |
| BLOCK_N | 32-128 | K/V block 行数 |
| num_warps | 4-8 | 每 block 的 warp 数 |
| num_stages | 1-2 | pipeline 阶段数 |
7.4 不同硬件的选择
# 根据 GPU 架构选择最优参数
def select_flash_config(gpu_arch, seq_len, d_head):
if gpu_arch == "sm80": # A100
if d_head <= 64:
return {"BLOCK_M": 128, "BLOCK_N": 128, "num_warps": 4}
elif d_head <= 128:
return {"BLOCK_M": 128, "BLOCK_N": 128, "num_warps": 8}
else: # d_head = 256
return {"BLOCK_M": 64, "BLOCK_N": 64, "num_warps": 8}
elif gpu_arch == "sm90": # H100
return {"BLOCK_M": 128, "BLOCK_N": 128, "num_warps": 4, "TMA": True}
八、超越 Flash Attention:Attenton 优化的探索方向
- Multi-Query Attention (MQA) / Grouped-Query Attention (GQA):减少 KV head 数,从根源降低 KV Cache 大小。Llama 2 70B 使用 GQA
- Sliding Window Attention (SWA):Mistral 采用,每个 token 只关注窗口内的前 W 个 token,复杂度降低到 O(N·W)
- Ring Attention:分布式场景下,将 KV 块分布在不同 GPU 上,通过 ring 通信实现超长序列的 Flash Attention
- Sparse Attention:通过局部敏感哈希(LSH)或学习到的稀疏模式,实现亚二次复杂度的 Attention
总结
Flash Attention 的成功不是偶然的工程 hack,而是从存储层次结构理论出发的系统性优化:
- 分块 (Tiling):将 O(N²) 分块适配 SRAM
- 在线统计量:通过修正因子实现单遍 softmax
- IO 复杂度下界:证明在 HBM-SRAM 模型下达到最优或接近最优
- 反向传播复用:通过重计算避免存储中间激活
理解 Flash Attention 的核心,不仅是学会调用一个 CUDA kernel,更是理解"算法设计必须考虑存储层次"这一现代系统优化的核心思想。

发表评论 取消回复