我第一次对 Flash Attention 起疑心,是在一次复现实验的时候。同一个模型、同一份数据,我只是把注意力算子从 PyTorch 的标准实现换成了 flash 版本。跑了一夜,第二天看日志,loss 曲线跟基准对不上了——差得不多,但足够让一个强迫症难受一天。

我的第一反应是:这玩意儿果然是有损的。
然后我翻到论文标题,差点噎住:FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness。Exact——精确。官方文档也反复强调,这是标准注意力的精确实现,不是近似。
一边是论文说 Exact,一边是我实测的差异。这两件事怎么同时成立?我花了一整个下午把它的 kernel 逻辑捋了一遍,才彻底想通。下面把过程讲给你听。
它没省计算,只是把矩阵拆开再算
标准注意力要算 softmax(QKT/√d)·V。麻烦在于 QKT 是一个 N×N 矩阵:序列长度 4096 时,它在 fp16 下要占 32MB;拉到 128k,这个数字变成 32GB,一张 A100 都塞不下。
Flash Attention 的思路是分块。把 Q、K、V 切成固定大小的小块,每次只算其中一块的局部注意力,把结果累加到输出块上,算完就扔——从头到尾,完整的 N×N 矩阵从不落地显存。
注意一个关键细节:公式一个都没改。该算的乘积、该做的 softmax、该加的累加,全部原样执行,只是换了顺序、换了位置。所以在实数域里,Flash Attention 算出的结果和标准注意力严格相等。这就是论文敢写 Exact 的原因。
"FlashAttention is an IO-aware exact attention algorithm that uses tiling to reduce the number of memory reads/writes between GPU high bandwidth memory (HBM) and GPU on-chip SRAM."
—— Tri Dao et al., FlashAttention (ICML 2022)
那为什么实际跑出来不一样?因为计算机压根不是在实数域里工作的。
罪魁祸首是 online softmax 里那步 rescale
Flash Attention 在每个分块上都要做一次局部 softmax。但完整的 softmax 需要先知道一整行的最大值和归一化常数——你眼前只有一个块,怎么知道整行的情况?
办法是边走边看:维护一个运行中的最大值 m,每处理一个块就更新它。这套技巧不是 Flash Attention 发明的,它来自 2018 年的一篇短论文,名字就叫 online softmax。
问题来了——假设第一个块算完,当前行最大值是 1.0。第二个块一算,最大值是 2.0。那之前那块已经算好的输出怎么办?整体缩放:把旧输出全部乘以 exp(1 − 2) ≈ 0.368。这一步有个专门的名字,叫 rescaling,它是 Flash Attention 与标准注意力数值差异的头号来源。
Flash Attention 的前向过程,概括起来是:
- 把 Q、K、V 切成固定大小的小块;
- 对每个 Q 块,依次遍历所有 K/V 块,计算局部分数 S = QiKjT/√d 和局部 softmax;
- 每当发现更大的行最大值 mnew,把旧输出和旧归一化常数乘以 exp(mold − mnew) 整体缩放;
- 把新分块的加权结果累加进去,最后除以累计的归一化常数。
每次 rescale 都是一次额外的浮点乘法,都会引入一次舍入。块数越多、序列越长,这种校正发生得越频繁,累积的偏差就越明显。
浮点数不认结合律,这口锅得 IEEE 754 背
但 rescale 只是表层。最底层的根源是:浮点数的加法不满足结合律。
你在数学课上学过 (a + b) + c = a + (b + c)。但在 IEEE 754 浮点数里,这句话不成立——每个数只有有限精度,加法的中间结果必须舍入,加法顺序变了,舍入的路径就变了,最后几位就是不一样。
标准注意力是先把一整行的 exp 全部算出来,统一求最大值、统一求和、统一归一化。Flash Attention 则是分块求局部 softmax,再用缩放因子把结果校正回全局视角。同一个数学表达式,两种完全不同的舍入轨迹,逐位一致才见鬼了。
所以这不是 bug,是所有浮点计算的宿命。你拿 PyTorch 和 TensorFlow 各自的标准注意力实现去对比,结果同样对不齐,原因完全相同。
反向传播也是一个道理。标准实现直接拿前向存好的 softmax 结果算梯度,Flash Attention 为了省内存,选择在反向时重新把 softmax 算一遍——重算的值跟前向的值在小数点后几位不一样,所以梯度也不逐位一致。内存和计算是省下来了,代价就是这笔永远对不齐的账。
差异有多大?1e-2 以内,前向不炸,训练会分叉
先给你一个直觉。bf16 的机器精度约 2⁻⁷ ≈ 0.008,fp16 约 2⁻¹⁰ ≈ 0.001。每次浮点运算的相对误差就在这个量级,Flash Attention 多了几次 rescale,误差再大也停留在 1e-2 以内。论文实测报告的最大差异同样在这个范围。
这里面还有两个重要的结论:
- 前向误差不会雪崩。每一层的归一化和非线性会把误差「揉平」,几十层叠加下来,相对偏差依然在 1e-2 到 1e-3 量级,不会爆炸。
- 但训练轨迹一定会分叉。1e-3 的扰动经过成千上万步优化器的迭代放大,两个实现最终会得到两组完全不同的权重。这跟换一个随机种子没有本质区别——最终指标在统计上没有差异,但逐位对齐就别想了。
想通这一点后,我对 Flash Attention 的「有损」彻底放心了:它的误差是随机舍入级的,没有系统性偏向。相比之下,Sparse Attention、低秩近似那些算法的误差是结构性的、有偏的,会随着网络变深累积成可感知的质量下降。这是本质区别。
什么时候这个「有损」会真的咬到你
两种场景需要小心。一是复现别人的实验——attention 后端不一致,loss 曲线就不可能重合;二是你自己做消融或对比实验——所有对照组必须用同一个 attention 后端,否则这 1e-3 的差异会混进你的结论里。
除此之外,我认为这个「有损」相当值得。它换来的是:
- 显存从 O(N²) 降到 O(N),序列越长差距越悬殊——这是长上下文模型能跑起来的物理前提;
- 论文实测 2~4 倍加速,FlashAttention-2 在此基础上又翻了约一倍;
- kernel 是确定性的:同一输入在同一 GPU 上每次结果一致。这一点反而比默认设置下的 cuBLAS 更令人安心。
回到开头的矛盾:「Exact」和「有差异」怎么共存?我的理解是:「Exact」说的是数学上的精确,「差异」来自浮点数的物理现实。论文没有否认后者,是很多人(包括我)一开始把「精确」误读成了「逐位一致」。
所以我的最终判断是:Flash Attention 的「有损」是一种诚实的取舍。你现在的 PyTorch 标准实现,携带的是同量级的舍入误差——你拿它当基准,不是因为它更精确,只是因为你更熟悉它。放弃对无损的执念,拥抱这种有损,你会发现它省下的显存和算力,已经撑起了大模型时代最长的那个上下文窗口。
原创文章,作者:guanweilu,如若转载,请注明出处:https://guanweilu.cn/article/547.html