FlashAttention:为何不用存完整注意力矩阵
注意力计算很容易产生一个N×N的中间矩阵:序列越长,保存它的代价越大。FlashAttention的出发点是改变数据搬运和中间结果的保存方式,用分块计算得到同一个稠密注意力定义,而不是默认删掉一部分token关系。
问题先出在中间矩阵
通常把计算写成 O = softmax(QKᵀ / √d)V,N是序列长度,d是头维度。如果把每个头完整的分数矩阵存下来,就需要N²个元素。N=8192、每元素2字节时,仅一个这样的矩阵就占128 MiB;这不是模型总显存估算。
完整N×N分数与概率
局部得分 → 更新统计量
不保存完整概率矩阵
FlashAttention原始论文于2022年公开,方法将输入分块并利用片上存储,减少GPU高带宽内存的读写。它还在反向过程中重算部分中间量。这里聚焦前向的数学机制,不复现论文训练成绩。
为什么不能平均各块的softmax
每块分别归一化,再把结果平均,并不等于全局softmax,因为各块的分母不同。要正确合并,必须保存它们相对于同一个最大值的指数总和。
对已经处理的分数,保存最大值m、归一化前总量l和加权和u。新分数a及对应标量值v到来时,令m′=max(m,a),旧统计量乘以exp(m−m′),新项权重为exp(a−m′)。于是:
l′ = exp(m − m′) l + exp(a − m′)
u′ = exp(m − m′) u + exp(a − m′) v
output = u′ / l′
减最大值使指数不超过1,避免直接对很大的正数求指数。矢量值V的情况,对u的各分量采用同一缩放因子。逐项递推是便于理解的最小例子;真正内核会处理块与矩阵乘法。
用三个分数检查结果
人工构造分数[0,1,2]与值[10,20,30]。下面两条计算路径都得到约25.752104:一条保存完整权重,另一条在线更新统计量。给全部分数加1000仍保持结果,这是softmax对统一平移不变的检验。
import math
def online(scores, values):
m, total, weighted = -math.inf, 0.0, 0.0
for score, value in zip(scores, values):
new_m = max(m, score)
old_scale = math.exp(m-new_m)
weight = math.exp(score-new_m)
total = old_scale*total + weight
weighted = old_scale*weighted + weight*value
m = new_m
return weighted / total
def reference(scores, values):
m = max(scores)
weights = [math.exp(s-m) for s in scores]
return sum(w*v for w,v in zip(weights,values))/sum(weights)
scores, values = [0.0,1.0,2.0], [10.0,20.0,30.0]
a, b = online(scores,values), reference(scores,values)
assert math.isclose(a,b,rel_tol=1e-12)
assert math.isclose(online([s+1000 for s in scores],values),b,
rel_tol=1e-12)
print(round(a,6))
代码已执行,并额外核对390组随机输入、长度1至39的结果。它只验证未加掩码、无dropout的一行标量加权和,不是FlashAttention内核,也没有验证反向梯度、GPU时间或完整模型。
“精确注意力”不等于逐位相同
这里的精确,是仍计算原来的稠密注意力数学定义。改变浮点运算顺序可能引入数值差异,不能要求任何精度和内核都逐位相同。掩码、dropout、头维度与硬件支持也必须按具体实现检查。
对于一般稠密注意力,各token两两关系仍然存在。不能把少存中间矩阵说成把计算量变为线性。选型时应同时记录内存峰值与延迟,并保持输入形状、精度和掩码一致。
它与PagedAttention也有不同侧重点:后者文章讨论KV缓存的组织;这篇讨论注意力算子内部的数据流。两项优化可以作用于不同层面,名称相似并不意味着可互换。
资料来源
FlashAttention原始预印本(2022);作者官方实现。安装与硬件支持应以实际选用版本为准。
补发说明:北京时间2026年10月10日实际补发,文章日期保留原计划时段。
支持
如果这篇文章对你有帮助,欢迎支持本站。
二维码可点击放大。更多支持方式见支持页面。


