FlashAttention:注意力的瓶颈不是算力,是内存带宽

Transformer 的 Self-Attention 计算理论上需要 O(N²) 的算力和内存,其中 N 是序列长度。很多人以为瓶颈是 O(N²) 的算力——计算量太大。但 FlashAttention 的作者发现事情不是这样的:标准 Attention 的真正瓶颈是内存带宽,不是算力。GPU 的 SRAM(片上缓存)读写速度比 HBM(显存)快几十倍,而标准实现在 HBM 上做了大量不必要的读写。

FlashAttention 的核心贡献:用分块(Tiling)和在线 Softmax 技术,把 HBM 访问次数从 O(N²) 降到 O(N),在数学上完全等价的前提下,训练提速 2-4 倍,同时支持更长的上下文窗口。

一、标准 Attention 的 IO 问题

先看标准 Attention 的计算流程:

# 标准 Attention(简化版)
# Q, K, V: shape [N, d],N 是序列长度,d 是 head dimension

S = Q @ K.T / sqrt(d)   # [N, N] 注意力分数矩阵
P = softmax(S)           # [N, N] 注意力权重
O = P @ V               # [N, d] 输出

问题在 SP 这两个 N×N 矩阵上。

以序列长度 N=4096、FP16 为例,S 矩阵大小是 4096×4096×2 bytes ≈ 32MB。这个矩阵必须写入 HBM,然后读出来计算 Softmax,写回 HBM,再读出来做矩阵乘法。标准实现里,这个 N×N 矩阵在 HBM 上来回读写了多次。

标准 Attention:N×N 矩阵在 HBM 多次读写 Q,K,V HBM S = QKᵀ/√d 写入 HBM N×N 矩阵 softmax(S) 读 HBM P 写 HBM N×N O = P·V 读 HBM 输出 O S 和 P 两个 N×N 矩阵在 HBM 反复读写 → IO 瓶颈 N=4096 时,S、P 矩阵各约 32MB,HBM IO 复杂度 O(N²)
标准 Attention:N×N 注意力矩阵在 HBM 多次读写,IO 复杂度 O(N²)

HBM 带宽大约是 SRAM 的 10-40 倍慢(见 GPU 资源模型那篇)。大量时间消耗在等待 HBM 传输数据上,GPU 的计算单元却在空转——这就是为什么标准 Attention 是内存带宽受限,而不是算力受限。

二、FlashAttention 的三个关键技术

1. 分块(Tiling):让数据留在 SRAM

FlashAttention 的基本思路:不要把整个 N×N 矩阵写入 HBM,而是把 Q、K、V 切成小块(Block),每次只把一个块加载进 SRAM,在片上完成计算,最后把结果写回 HBM。

如果 SRAM 能装下一个块的 Q、K、V 数据,整个计算过程就只需要:

  • 把 Q、K、V 从 HBM 读一遍:O(Nd)
  • 把输出 O 写回 HBM 一遍:O(Nd)
  • 合计 HBM 访问:O(Nd),而不是 O(N²)

问题来了:Softmax 需要一整行的数据才能计算归一化因子(分母 sum(exp(x_i))),但分块计算时,你每次只能看到一小块数据,怎么算完整行的 Softmax?

2. 在线 Softmax(Online Softmax):增量计算归一化

这是 FlashAttention 最精妙的部分。它利用数学上的分块等价性,在不看完整行的情况下,增量地维护一个"运行中的 Softmax 估计"。

标准的数值稳定 Softmax 是:

def softmax(x):
    m = max(x)              # 减去最大值保证数值稳定
    e = exp(x - m)
    return e / sum(e)

在线 Softmax 的思路:假设你先看到 x 的前半段(x₁),再看到后半段(x₂)。你可以维护两个运行状态——当前最大值 m归一化因子 ℓ,处理每个新块时更新它们:

def online_softmax_step(m_prev, l_prev, O_prev, Q_block, K_block, V_block):
    # 计算这个 K_block 对应的注意力分数
    S_block = Q_block @ K_block.T / sqrt(d)

    # 更新运行最大值
    m_new = max(m_prev, max(S_block))

    # 更新归一化因子(用新的最大值重新缩放旧的)
    l_new = exp(m_prev - m_new) * l_prev + sum(exp(S_block - m_new))

    # 更新输出(重新缩放旧的,加上新块的贡献)
    O_new = (exp(m_prev - m_new) * l_prev * O_prev
             + exp(S_block - m_new) @ V_block) / l_new

    return m_new, l_new, O_new

关键在于 exp(m_prev - m_new) 这个重缩放因子:每当发现新的更大值,它会把之前的贡献按比例调小,保证最终结果和一次性看完整行的 Softmax 在数学上完全等价。

这个技巧让 FlashAttention 可以按块处理,而不需要把整个 N×N 矩阵物化到 HBM 上。

3. 反向传播重计算(Recomputation):省内存,不省精度

训练时不只要前向传播,还要反向传播求梯度。标准 Attention 在前向时会把 P(N×N 的注意力权重矩阵)存储下来,供反向传播使用。这就是为什么序列越长,训练时的显存消耗越爆。

FlashAttention 的选择:前向传播时不保存 P,反向传播时重新计算它

需要保存的只有:

  • Q、K、V(原本就要保存的)
  • 每个块的 softmax 统计量 m 和 ℓ(每行只有两个标量)

反向传播时,用这些信息重新计算出需要的 P 块。多了一次重计算(额外的算力消耗),但避免了把 N×N 矩阵写入 HBM(大量节省 IO)。由于 IO 才是瓶颈,这个交换是划算的:实际测量表明重计算带来的算力开销远小于节省的 IO 开销。

这个思路和梯度检查点(Gradient Checkpointing)类似,但粒度更细,直接在 Attention 层内部实现。

三、FlashAttention 的完整计算图

FlashAttention:分块加载到 SRAM,在片上完成计算 HBM(慢) Q 块 K 块 V 块 O 块 写回一次 (N×N 矩阵 不存储) SRAM(快,片上) Q_i K_j V_j S_ij = Q_i·K_j/√d 在线 Softmax 更新 m, ℓ, O_i O_i 累积 循环所有 j(K/V 块),O_i 累积完成后写回 HBM → HBM 访问 O(Nd)
FlashAttention:Q/K/V 分块加载进 SRAM,片上完成计算,不把 N×N 矩阵写入 HBM

具体算法流程:

  1. 把 Q 按行分成若干块 Q₁, Q₂, …(外层循环)
  2. 把 K、V 按列分成若干块 K₁, K₂, … 和 V₁, V₂, …(内层循环)
  3. 对每个 (Q_i, K_j, V_j) 组合:加载进 SRAM,计算 S_ij = Q_i · K_j^T,用在线 Softmax 更新运行统计量 m 和 ℓ,累积输出 O_i
  4. 处理完所有 K/V 块后,O_i 计算完成,写回 HBM

整个过程,Q、K、V 各从 HBM 读一次,O 写回一次,HBM 访问总量是 O(Nd),而不是 O(N²)。

四、FlashAttention-2 的改进

FlashAttention-2 在 v1 基础上做了三处优化,整体再提速约 2 倍:

1. 减少非矩阵乘法运算

GPU 对矩阵乘法有专门的 Tensor Core 加速,吞吐量极高;但对其他运算(加法、乘法、exp)则没有这种加速,算力远低。这类运算叫做"非矩阵乘法 FLOPs"(non-matmul FLOPs)。

在 FlashAttention v1 中,每处理完一个 K/V 块,就要对累积的 O_i 做一次重缩放(乘以 exp(m_old - m_new) * l_old / l_new)。如果有 T 个 K/V 块,就有 T 次重缩放,这些都是非矩阵乘法运算。

FlashAttention-2 的改进:延迟重缩放。在内层循环中不做归一化,先累积一个未归一化的 O_i;等所有 K/V 块处理完后,只做一次最终归一化:

# v1:每个 K/V 块处理完都重缩放(T 次)
for j in range(T):
    S_ij = Q_i @ K_j.T / sqrt(d)
    m_new = max(m, rowmax(S_ij))
    O_i = (exp(m - m_new) * l * O_i + exp(S_ij - m_new) @ V_j)
    l = exp(m - m_new) * l + rowsum(exp(S_ij - m_new))
    m = m_new
O_i = O_i / l  # 最终归一化

# v2:累积未归一化的 O_tilde,最后一次归一化(1 次)
for j in range(T):
    S_ij = Q_i @ K_j.T / sqrt(d)
    m_new = max(m, rowmax(S_ij))
    O_tilde = exp(m - m_new) * O_tilde + exp(S_ij - m_new) @ V_j
    l = exp(m - m_new) * l + rowsum(exp(S_ij - m_new))
    m = m_new
O_i = O_tilde / l  # 只做一次归一化

非矩阵乘法运算从 T 次降到 1 次,GPU 利用率提升明显。

2. 序列维度并行

FlashAttention v1 只在 batch 和 head 维度上并行(每个 head 分配给一个 thread block)。当 batch size 很小或 head 数很少时(比如推理时 batch=1,head=8),GPU 利用率很低,大量 CUDA core 空闲。

v2 增加了在序列维度上的并行:同一个 head 的不同 Q 块可以分配给不同的 thread block 并行计算。这让序列很长但 batch 很小的场景(推理时的典型场景)也能充分利用 GPU。

3. Warp 分工优化

在一个 thread block 内部,v1 把不同的 Q 块分给不同的 warp(32 线程组),warp 之间需要同步共享 K/V 数据——同步有开销,而且 K/V 数据会被多个 warp 重复读取。

v2 改为:同一个 Q 块被一个 warp 处理,把不同的 K/V 块分给不同 warp 处理(然后 reduce 合并结果)。这减少了 warp 间同步,也让 K/V 数据的读取更高效。

五、实际效果和应用

速度提升:在 A100 上,FlashAttention v1 比标准 Attention 快约 2-4 倍(序列越长提升越明显);FlashAttention-2 比 v1 再快约 2 倍,接近理论峰值带宽利用率。

内存节省:不再存储 N×N 注意力矩阵,训练时的峰值显存从 O(N²) 降到 O(N)。这直接让更长的上下文成为可能——GPT-3 用标准 Attention 最多支持 2K 上下文,而 FlashAttention 让训练和推理百 K 级别的上下文成为可行的工程方案。

精度无损:FlashAttention 在数学上与标准 Attention 完全等价,不是近似算法,没有任何精度损失。

已成标准**:PyTorch 2.0+ 的 F.scaled_dot_product_attention 在支持的设备上自动使用 FlashAttention 内核;xFormers、DeepSpeed、Megatron-LM、HuggingFace Transformers 等主流框架都已集成。现在用 Transformer 训练或推理,几乎必然在用 FlashAttention。

import torch
import torch.nn.functional as F

# PyTorch 2.0+ 自动使用 FlashAttention(如果硬件支持)
output = F.scaled_dot_product_attention(
    query, key, value,
    attn_mask=None,
    dropout_p=0.0,
    is_causal=True   # 因果掩码(用于自回归生成)
)

六、总结

FlashAttention 系列的核心洞察是:

  • Attention 的瓶颈是内存带宽,不是算力——N×N 矩阵在 HBM 上反复读写,导致 GPU 大量时间在等数据
  • 分块(Tiling):把 Q/K/V 切成小块,让计算留在快速的 SRAM 里,避免把 N×N 矩阵写入 HBM,HBM 访问从 O(N²) 降到 O(N)
  • 在线 Softmax:通过维护运行最大值和归一化因子,在不看完整行的情况下增量计算 Softmax,使分块成为可能
  • 反向重计算:不保存 N×N 的注意力权重矩阵,反向传播时重新计算,大幅节省训练显存
  • FlashAttention-2:延迟归一化减少非矩阵乘法运算、增加序列维度并行、优化 warp 分工,在 v1 基础上再快约 2 倍

它是一个很好的例子:算法上没有做任何近似,结果完全一样,但通过深刻理解硬件的 IO 特性,在工程实现上带来了数倍的性能提升。