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] 输出
问题在 S 和 P 这两个 N×N 矩阵上。
以序列长度 N=4096、FP16 为例,S 矩阵大小是 4096×4096×2 bytes ≈ 32MB。这个矩阵必须写入 HBM,然后读出来计算 Softmax,写回 HBM,再读出来做矩阵乘法。标准实现里,这个 N×N 矩阵在 HBM 上来回读写了多次。
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 的完整计算图
具体算法流程:
- 把 Q 按行分成若干块 Q₁, Q₂, …(外层循环)
- 把 K、V 按列分成若干块 K₁, K₂, … 和 V₁, V₂, …(内层循环)
- 对每个 (Q_i, K_j, V_j) 组合:加载进 SRAM,计算 S_ij = Q_i · K_j^T,用在线 Softmax 更新运行统计量 m 和 ℓ,累积输出 O_i
- 处理完所有 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 特性,在工程实现上带来了数倍的性能提升。