FlashAttention:为何不用存完整注意力矩阵

黎 浩然/ 9 10 月, 2026/ 大语言模型/LARGELANGUAGEMODEL/LLM, 机器学习/MACHINELEARNING, 研究生/POSTGRADUATE, 计算机/COMPUTER/ 0 comments

注意力计算很容易产生一个N×N的中间矩阵:序列越长,保存它的代价越大。FlashAttention的出发点是改变数据搬运和中间结果的保存方式,用分块计算得到同一个稠密注意力定义,而不是默认删掉一部分token关系。

文章目录
  1. 问题先出在中间矩阵
  2. 为什么不能平均各块的softmax
  3. 用三个分数检查结果
  4. “精确注意力”不等于逐位相同
  5. 资料来源

问题先出在中间矩阵

通常把计算写成 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缓存的组织;这篇讨论注意力算子内部的数据流。两项优化可以作用于不同层面,名称相似并不意味着可互换。

English version

资料来源

FlashAttention原始预印本(2022);作者官方实现。安装与硬件支持应以实际选用版本为准。

补发说明:北京时间2026年10月10日实际补发,文章日期保留原计划时段。

支持

如果这篇文章对你有帮助,欢迎支持本站。

微信支持二维码,点击查看大图
微信
支付宝支持二维码,点击查看大图
支付宝
Buy Me a Coffee,点击支持本站
Buy Me a Coffee

二维码可点击放大。更多支持方式见支持页面。

Share this Post

Leave a Comment

您的邮箱地址不会被公开。 必填项已用 * 标注

*
*