Transformer 的自注意力在数学上并不复杂,但当序列变长时,它会突然变得又慢又吃显存。FlashAttention 的洞察是:真正的瓶颈不在浮点运算,而在显存的读写。把这一点想清楚,优化方向就完全变了。

标准注意力为什么慢

标准注意力分三步:先算打分矩阵 S = Q·Kᵀ,再做 P = softmax(S),最后 O = P·V。问题出在中间那个 S:它的形状是 N×NN 是序列长度。序列翻倍,这个矩阵就变成四倍大。

更要命的是,常规实现会把 SP 这两个 N×N 矩阵完整写回显存(HBM),再读回来做下一步。于是注意力的时间几乎全花在搬运这两个大矩阵上——它是访存受限(memory-bound),而不是算力受限。GPU 的计算单元大量空转,等着数据从显存里挪进挪出。

标准注意力反复把 N×N 中间矩阵写回 HBM,而 FlashAttention 在 SRAM 内分块计算,N×N 从不落地

图 1:标准注意力把 N×N 矩阵反复读写显存;FlashAttention 分块后让中间结果只在片上 SRAM 里流动。

IO 感知:把计算搬进 SRAM

GPU 的存储是分层的:片上的 SRAM 很小(每个计算单元只有几十 KB),但极快;显存 HBM 很大(几十 GB),却慢一个数量级。

FlashAttention 论文给出的 A100 数据里,SRAM 的带宽是 HBM 的十几倍。这个倍数来自论文,不是我在本机实测的结果——不同硬件差别很大,这里只用来说明「层级之间存在数量级差距」这个定性事实。

既然差距这么大,正确的做法就是尽量让数据待在 SRAM 里、少碰 HBM。FlashAttention 把 QKV 切成小块,每次只把能塞进 SRAM 的一小块搬上片,在片上把这一块的注意力算完,N×N 的中间矩阵从头到尾都不在 HBM 里落地

在线 softmax:分块还能保持精确

分块有个坎:softmax 需要一整行的全局最大值来做数值稳定,可我们是一块一块处理的,处理前面的块时并不知道后面块里会不会冒出更大的值。

FlashAttention 用在线 softmax跨过这道坎——维护三个运行统计量:当前行最大值 m、指数和 l、以及输出累加器 O。每来一个新块,就地更新:

S_ij = Q_i · K_jᵀ
m'   = max(m, rowmax(S_ij))          # 刷新运行最大值
P̃    = exp(S_ij − m')
l    ← e^(m−m') · l + rowsum(P̃)      # 旧的和按新最大值重新缩放
O    ← e^(m−m') · O + P̃ · V_j        # 已累加的输出也一起缩放
m    ← m'

关键在那个 e^(m−m') 因子:一旦新块带来了更大的最大值,就用它把已经累加进去的 Ol 追溯性地重新缩放一遍。收尾时再 O ← O / l 归一化。这样得到的结果和「先看到整行、再一次性做 softmax」逐位相等

固定一个 Q 块,遍历 K/V 块;每块更新运行最大值 m、指数和 l 与输出累加器 O,新块带来更大最大值时用 e^(m−m') 重新缩放

图 2:在线 softmax 的状态更新。分块是增量的,运行统计量保证最终结果与一次性 softmax 完全等价。

这一点很重要:FlashAttention 不是近似。它和稀疏注意力、低秩近似那类方法有本质区别——后者为了省显存牺牲了精度,而 FlashAttention 给出的是精确解,只是把计算重新组织了一遍。

反向传播用重算换显存

前向不落地 N×N 矩阵,反向传播时却需要 P 来求梯度,怎么办?答案是重算:反向时不去读存好的 P(根本没存),而是用 QKV 的块,配合前向留下的 ml 这几个很小的统计量,把需要的那块 P 现场重新算出来。

这是一次经典的取舍——多花一点算力,换回大量显存。因为注意力本就是访存受限的,多出来的这点浮点运算几乎不影响墙钟时间,省下的显存却让更长的上下文成为可能。

后续版本做了什么

版本主要改进方向
FlashAttention-2更好的并行与工作划分:减少非矩阵乘的开销、在序列维度上并行、优化 warp 间的任务分配
FlashAttention-3面向 Hopper 架构:利用异步指令与更低精度(FP8),进一步压榨新硬件的吞吐

两代的思路一脉相承:继续减少不必要的数据搬运和调度开销,并贴着新硬件的特性去重排计算。具体的加速倍数各家测法不一,这里就不引用了——它们高度依赖 GPU 型号、序列长度和精度设置。

什么时候值得上

长上下文的训练、以及推理时的 prefill 阶段,是 FlashAttention 收益最明显的地方:序列越长,N×N 的访存成本增长得越快,省下来的也越多。今天主流训练框架和推理引擎基本都已内置它,多数时候你不需要自己实现,但理解它「用访存视角而非算力视角看问题」的思路,对做任何 GPU 性能优化都有借鉴意义。