Transformer 的自注意力在数学上并不复杂,但当序列变长时,它会突然变得又慢又吃显存。FlashAttention 的洞察是:真正的瓶颈不在浮点运算,而在显存的读写。把这一点想清楚,优化方向就完全变了。
标准注意力为什么慢
标准注意力分三步:先算打分矩阵 S = Q·Kᵀ,再做 P = softmax(S),最后 O = P·V。问题出在中间那个 S:它的形状是 N×N,N 是序列长度。序列翻倍,这个矩阵就变成四倍大。
更要命的是,常规实现会把 S 和 P 这两个 N×N 矩阵完整写回显存(HBM),再读回来做下一步。于是注意力的时间几乎全花在搬运这两个大矩阵上——它是访存受限(memory-bound),而不是算力受限。GPU 的计算单元大量空转,等着数据从显存里挪进挪出。
IO 感知:把计算搬进 SRAM
GPU 的存储是分层的:片上的 SRAM 很小(每个计算单元只有几十 KB),但极快;显存 HBM 很大(几十 GB),却慢一个数量级。
FlashAttention 论文给出的 A100 数据里,SRAM 的带宽是 HBM 的十几倍。这个倍数来自论文,不是我在本机实测的结果——不同硬件差别很大,这里只用来说明「层级之间存在数量级差距」这个定性事实。
既然差距这么大,正确的做法就是尽量让数据待在 SRAM 里、少碰 HBM。FlashAttention 把 Q、K、V 切成小块,每次只把能塞进 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') 因子:一旦新块带来了更大的最大值,就用它把已经累加进去的 O 和 l 追溯性地重新缩放一遍。收尾时再 O ← O / l 归一化。这样得到的结果和「先看到整行、再一次性做 softmax」逐位相等。
这一点很重要:FlashAttention 不是近似。它和稀疏注意力、低秩近似那类方法有本质区别——后者为了省显存牺牲了精度,而 FlashAttention 给出的是精确解,只是把计算重新组织了一遍。
反向传播用重算换显存
前向不落地 N×N 矩阵,反向传播时却需要 P 来求梯度,怎么办?答案是重算:反向时不去读存好的 P(根本没存),而是用 Q、K、V 的块,配合前向留下的 m、l 这几个很小的统计量,把需要的那块 P 现场重新算出来。
这是一次经典的取舍——多花一点算力,换回大量显存。因为注意力本就是访存受限的,多出来的这点浮点运算几乎不影响墙钟时间,省下的显存却让更长的上下文成为可能。
后续版本做了什么
| 版本 | 主要改进方向 |
|---|---|
| FlashAttention-2 | 更好的并行与工作划分:减少非矩阵乘的开销、在序列维度上并行、优化 warp 间的任务分配 |
| FlashAttention-3 | 面向 Hopper 架构:利用异步指令与更低精度(FP8),进一步压榨新硬件的吞吐 |
两代的思路一脉相承:继续减少不必要的数据搬运和调度开销,并贴着新硬件的特性去重排计算。具体的加速倍数各家测法不一,这里就不引用了——它们高度依赖 GPU 型号、序列长度和精度设置。
什么时候值得上
长上下文的训练、以及推理时的 prefill 阶段,是 FlashAttention 收益最明显的地方:序列越长,N×N 的访存成本增长得越快,省下来的也越多。今天主流训练框架和推理引擎基本都已内置它,多数时候你不需要自己实现,但理解它「用访存视角而非算力视角看问题」的思路,对做任何 GPU 性能优化都有借鉴意义。
讨论
评论
还没有评论,来留下第一条讨论吧。