从第一性原理推导 Flash Attention

从第一性原理推导 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 内部:

  1. 将 Q_i 的 B_r 行分配到多个 warp
  2. 每个 warp 负责若干行的 softmax 计算
  3. Warp 之间通过 warp shuffle 指令交换 m 和 l 的归约结果
  4. 使用 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

关键工程决策:

  1. 不存储 P 矩阵:前向时维护在线统计量,反向时重新计算 S_ij 和 softmax。牺牲计算换取内存,净收益为正(因为瓶颈在 IO 不在计算)
  2. dQ 需要完整的前向遍历:每个 dQ_i 块需要所有 K_j 的信息,因此 dQ 的计算外层循环在 j,内层在 i
  3. dK, dV 可以与 dQ 合并计算:通过精心设计的 loop tiling 实现单次前向遍历

六、Flash Attention 2 & 3 的关键改进

6.1 Flash Attention 2:减少非矩阵乘法运算

FA1 的改进方向是 IO 复杂度,FA2 进一步优化计算效率:

  1. 减少非 FMA 操作:softmax 中的 exp 和除法在 CUDA 中开销大,FA2 通过更优的 warp 分工减少这类操作
  2. 降低 block 间的同步开销:FA1 中 Q 的外层循环导致每处理一个新 Q block 需要重新加载所有 KV;FA2 将循环改为 K/V 外层 + Q 内层,减少同步
  3. 更好的 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 指令:

  1. TMA 异步拷贝:硬件加速的 HBM↔SRAM DMA,无需显式 shared memory 管理
  2. WGMMA 指令:直接在 shared memory 上执行矩阵乘法,减少寄存器↔SRAM 的搬运
  3. 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 优化的探索方向

  1. Multi-Query Attention (MQA) / Grouped-Query Attention (GQA):减少 KV head 数,从根源降低 KV Cache 大小。Llama 2 70B 使用 GQA
  2. Sliding Window Attention (SWA):Mistral 采用,每个 token 只关注窗口内的前 W 个 token,复杂度降低到 O(N·W)
  3. Ring Attention:分布式场景下,将 KV 块分布在不同 GPU 上,通过 ring 通信实现超长序列的 Flash Attention
  4. Sparse Attention:通过局部敏感哈希(LSH)或学习到的稀疏模式,实现亚二次复杂度的 Attention

总结

Flash Attention 的成功不是偶然的工程 hack,而是从存储层次结构理论出发的系统性优化:

  1. 分块 (Tiling):将 O(N²) 分块适配 SRAM
  2. 在线统计量:通过修正因子实现单遍 softmax
  3. IO 复杂度下界:证明在 HBM-SRAM 模型下达到最优或接近最优
  4. 反向传播复用:通过重计算避免存储中间激活

理解 Flash Attention 的核心,不仅是学会调用一个 CUDA kernel,更是理解"算法设计必须考虑存储层次"这一现代系统优化的核心思想。

点赞(0) 打赏

评论列表 共有 0 条评论

暂无评论
立即
投稿

微信公众账号

微信扫一扫加关注

发表
评论
返回
顶部