上一篇:解析 FlashAttention(1):从标准 Attention 讲起
参考:万字长文详解FlashAttention v1/v2
整个 FlashAttention v1 的分块运算可视化过程: flash_attention_visualization.html
1. 背景 & 动机
1.1 标准 Attention 的核心矛盾
对于标准 Attention,公式如下:
Attention(Q,K,V)=softmax(dQK⊤)V
它具体的计算过程为:
S=QK⊤→P=softmax(S)→O=PV
涉及的内存操作如下,它主要使用了 HBM:

图中一共包含八次 HBM 的矩阵读写操作。这八次读写操作分别为:
- Line 1: 对 Q, K 的读取共两次,对 S 的写入一次,读写总共三次;
- Line 2: 对 S 读取一次,对 P 写入一次,读写总共两次;
- Line 3: 对 P, V 的读取共两次,对 O 的写入一次,读写总共三次。
1.2 FlashAttention 的核心思路
FlashAttention 的动机建立在三个观察上:
| 观察 |
含义 |
| 1. 内存层级差异巨大 |
GPU 的 SRAM(如 Shared Memory)比 HBM 快约 10–100 倍,但容量极小(如 A100 每 SM 仅 192KB)。标准 Attention 无视这一层级,把中间矩阵 S,P 放在 HBM 中反复读写,导致算法受限于 HBM 带宽,成为 memory-bound。 |
| 2. Attention 本不需要 O(N2) 显存 |
最终输出 O 只有 N×d,理论额外内存可以是 O(N)。存储巨大的 S,P 是实现方式的浪费,不是算法必需。 |
| 3. Softmax 可以"在线"算 |
以前认为 softmax 必须看到全部数字才能归一化。但实际上可以通过维护一个滑动最大值和累加和,在流式看到数据时增量地得到正确结果(Online Softmax)。 |
具体怎么做:
(1)Tiling(分块)
把 Q,K,V 切成足够小的块,使得一小块 Qi 和 Kj 能在 SRAM 里完成矩阵乘法。这样 Sij 的每个块在 SRAM 里生成,用完即弃,从不写回 HBM。
(2)Online Softmax(在线归一化)
Softmax 需要全局信息(最大值、指数和),分块计算时怎么办?
技巧是维护两个统计量:
- m:当前见过的最大值(用于数值稳定性)
- ℓ:当前指数和
当新的数据块进来时,用这两个量修正之前的结果,逐步逼近全局 softmax。这样不需要等所有 S 算完就能开始归一化。
(3)反向传播的重计算(Recomputation)
训练时反向传播需要 P 的梯度。既然 P 没存,FlashAttention 在反向时重新计算 P——但它只需要重新加载 Q,K 的小块,在 SRAM 里快速重算,成本远低于存储和读取巨大的 P 矩阵。
下面正式开始介绍 FlashAttention v1 版本的实现。
2. 前向传播: 标准 softmax 及其数值稳定版本
2.1 标准 softmax 数值稳定版本
以一维向量举例,标准 softmax 对向量 x∈RB 的第 i 个分量计算如下:
softmax(xi)=∑j=1Bexjexi(1)
由于分子和分母均包含指数项,当 xi 较大时 exi 容易上溢,当 x 中各元素均为较大负值时 exi 容易下溢导致分母为零。因此引入稳定版本。设 m(x)=maxjxj,在公式 (1) 的分子和分母上同乘 e−m(x),数学等价:
softmax(xi)=∑j=1Bexj⋅e−m(x)exi⋅e−m(x)=∑j=1Bexj−m(x)exi−m(x)(2)
基于该等价形式,定义平移后的指数向量:
f(x)=[ex1−m(x), ex2−m(x), …, exB−m(x)](3)
再定义 EXP 求和项:
l(x)=j=1∑Bf(x)j=j=1∑Bexj−m(x)(4)
最终稳定版 softmax 写为:
softmax(x)=l(x)f(x)(5)
此版本中分子最大项为 exmax−m(x)=e0=1,分母至少为 1,彻底消除了数值溢出的风险。
2.2 分块策略
设待计算 softmax 的向量 x∈R2B,将其按列一切为二:
x=[x(1), x(2)](6)
其中 x(1)∈RB 为第 1 个分块,x(2)∈RB 为第 2 个分块。上标 (1) 与 (2) 表示分块序号,下标 1,2,…,B 表示该分块内的元素索引。目标是先处理 x(1),再处理 x(2),最终得到与全局一次性计算完全一致的 softmax 结果。
2.3 第一块的局部 softmax 与全局统计量初始化
处理 x(1),计算其局部统计量。局部最大值:
m(x(1))=j=1maxBxj(1)(7)
局部平移指数向量:
f(x(1))=[ex1(1)−m(x(1)), ex2(1)−m(x(1)), …, exB(1)−m(x(1))](8)
局部 EXP 求和项:
l(x(1))=j=1∑Bf(x(1))j=j=1∑Bexj(1)−m(x(1))(9)
局部 softmax:
softmax(x(1))=l(x(1))f(x(1))(10)
此时 softmax(x(1)) 是局部的而非全局的,原因有二:其一,分子减去的最大值是 m(x(1)) 而非全局 m(x),导致平移基准不一致;其二,分母是 l(x(1)) 而非全局求和项,导致归一化因子只是局部和而非全体和。处理完 x(1) 后,初始化两个全局标量。当前全局最大值:
mmax=m(x(1))(11)
当前全局 EXP 求和项:
lall=l(x(1))(12)
2.4 第二块的局部 softmax
处理 x(2),同样计算局部统计量。局部最大值:
m(x(2))=j=1maxBxj(2)(13)
局部平移指数向量:
f(x(2))=[ex1(2)−m(x(2)), ex2(2)−m(x(2)), …, exB(2)−m(x(2))](14)
局部 EXP 求和项:
l(x(2))=j=1∑Bf(x(2))j=j=1∑Bexj(2)−m(x(2))(15)
局部 softmax:
softmax(x(2))=l(x(2))f(x(2))(16)
2.5 更新全局统计量
处理完 x(2) 后,需要利用 x(2) 的信息更新此前保存的两个全局标量 mmax 与 lall,以便后续将各局部 softmax 合并为全局 softmax。
更新全局最大值:
mmaxnew=max(mmax, m(x(2)))(17)
其含义为更新后的全局最大值是此前全局最大值与当前分块最大值中较大的那一个。
接下来更新全局 EXP 求和项。目标是得到以新的全局最大值 mmaxnew 为基准的全体指数和。对于此前全局求和项 lall,其当前以 mmax 为指数基准,即:
lall=k=1∑Bexk(1)−mmax
要将其基准由 mmax 调整至 mmaxnew,需对每一项同乘 emmax−mmaxnew:
lall⋅emmax−mmaxnew=k=1∑Bexk(1)−mmax⋅emmax−mmaxnew=k=1∑Bexk(1)−mmaxnew(18)
同理,对于当前分块 x(2) 的局部求和项 l(x(2)),其当前以 m(x(2)) 为基准:
l(x(2))=k=1∑Bexk(2)−m(x(2))
要将其基准调整至 mmaxnew,需同乘 em(x(2))−mmaxnew:
l(x(2))⋅em(x(2))−mmaxnew=k=1∑Bexk(2)−m(x(2))⋅em(x(2))−mmaxnew=k=1∑Bexk(2)−mmaxnew(19)
将调整后的两部分求和相加,即得到以 mmaxnew 为基准的全体指数和:
lallnew=emmax−mmaxnew⋅lall+em(x(2))−mmaxnew⋅l(x(2))(20)
为理解第二项 em(x(2))−mmaxnew⋅l(x(2)) 的来源,先将 l(x(2)) 展开:
l(x(2))=k=1∑Bexk(2)−m(x(2))(21)
将其变换为以新的全局最大值 mmaxnew 为基准:
lnew(x(2))=l(x(2))⋅em(x(2))−mmaxnew=k=1∑Bexk(2)−m(x(2))⋅em(x(2))−mmaxnew=k=1∑Bexk(2)−m(x(2))+m(x(2))−mmaxnew=k=1∑Bexk(2)−mmaxnew(22)
此时 l(x(2)) 更新为全局的。也就是说,通过对 l(x(2)) 乘上额外的项 em(x(2))−mmaxnew 即可把 l(x(2)) 更新为全局的 lnew(x(2))。简而言之,当需要把某个 EXP 求和项 l 更新为全局的时,只要将其乘以 em−mmaxnew 即可,其中 m 表示当前 l 对应的最大值,mmaxnew 表示当前全局最大值。
回到公式 (20),lall 对应的最大值是 mmax,当前全局最大值是 mmaxnew,所以可以乘以项 emmax−mmaxnew 来更新 lall(参考公式 (20) 等式右方的第一项)。同理再使用 em(x(2))−mmaxnew 来更新 l(x(2))(参考公式 (20) 等式右方的第二项)。最后将更新后的两项求和得到当前的 EXP 求和项 lallnew。
2.6 将局部 softmax 更新为全局
为什么要将局部 softmax 更新为全局?因为 softmax 的分母必须是对整个向量 x(包含所有分块)的指数求和,分子必须是每个元素相对于全局最大值的指数。如果不更新为全局,那么每个分块内的 softmax 值只是基于该分块内部归一化的,无法反映元素在整个向量中的相对权重。因此,在处理完当前分块后,必须将此前各分块与当前分块的 softmax 值都重新归一化到统一的全局基准上。
基于上述更新 l 的方法,也能直接更新 softmax 值。参考公式 (16),可知当前的分子和分母都是局部的,所以需要将它们分别更新至全局。
先看分子部分 f(x(2))。f(x(2)) 由公式 (14) 定义,可将其做如下更新:
fnew(x(2))=f(x(2))⋅em(x(2))−mmaxnew=[ex1(2)−m(x(2)), …, exB(2)−m(x(2))]⋅em(x(2))−mmaxnew=[ex1(2)−mmaxnew, …, exB(2)−mmaxnew](23)
此时 fnew(x(2)) 中的每一项都是全局的。很容易发现,更新 f(x(2)) 的方法与更新 l(x(2)) 其实是一样的,都是乘以项 em(x(2))−mmaxnew。
基于公式 (23) 的结果,可以首先将公式 (16) 的分子乘以 em(x(2))−mmaxnew 来将分子更新为全局的:
softmaxtemp(x(2))=softmax(x(2))⋅em(x(2))−mmaxnew=l(x(2))f(x(2))⋅em(x(2))−mmaxnew=l(x(2))fnew(x(2))(24)
注意公式 (24) 中的 softmaxtemp(x(2)) 仅仅是更新了分子,还没有更新分母,所以它不是最终结果。softmaxtemp(x(2)) 离最终结果只有一步之差:分母中的局部 EXP 求和项 l(x(2)) 需要被替换成全局 EXP 求和项 lallnew,而 lallnew 已经在公式 (20) 中计算出来了。
将分母替换为全局求和项 lallnew,得到最终全局 softmax:
softmaxnew(x(2))=softmaxtemp(x(2))⋅lallnewl(x(2))=lallnewfnew(x(2))(25)
将上述步骤合并,直接写出从局部 softmax 到全局 softmax 的更新式:
softmax(new)(x(2))=lallnewsoftmax(x(2))⋅l(x(2))⋅em(x(2))−mmaxnew(26)
同理,对 x(1) 也执行相同的全局化更新。x(1) 的分子 f(x(1)) 当前以 mmax 为基准,需调整至 mmaxnew:
fnew(x(1))k=f(x(1))k⋅emmax−mmaxnew=exk(1)−mmax⋅emmax−mmaxnew=exk(1)−mmaxnew
因此:
softmax(new)(x(1))=lallnewsoftmax(x(1))⋅l(x(1))⋅emmax−mmaxnew(27)
把 softmax(new)(x(1)) 和 softmax(new)(x(2)) 直接拼接,就是整个向量 x 的 softmax。
所有更新均不需要重新访问 x(1) 或 x(2) 的原始向量值,仅需之前保存的局部统计量与全局统计量。
2.7 全局统计量最终赋值
在当前的示例中,待计算 softmax 的向量为 x∈R2B={x(1),x(2)},所以此时把 softmax(new)(x(1)) 和 softmax(new)(x(2)) 直接拼接,就是整个向量 x 的 softmax。
但是,如果切分的块更多,那么处理完当前分块后,将需要新的全局统计量赋值给全局变量,为下一分块做准备:
mmax=mmaxnew(28)
以及:
lall=lallnew(29)
这里的一维向量 x ,放到 FlashAttention 里就是注意力分数矩阵 S=QKT 的某一行。
3. Attention 输出的分块增量更新
3.1 问题设定
设单头维度为 N=2,特征维度为 d。输入矩阵按行切分为 Br=Bc=1 的块:
Q=[q1q2]∈R2×d,K=[k1k2]∈R2×d,V=[v1v2]∈R2×d
其中 qi,kj,vj∈R1×d 均为行向量。记注意力分数:
sij=qikj⊤∈R
为第 i 个 query 向量与第 j 个 key 向量的内积。
对于输出矩阵 O∈R2×d,其每一行都是对应 query 向量与所有 key 向量和 value 向量的 Attention 结果。以下以第 1 行 o1(对应 q1)为例,展示其输出如何通过分块增量方式逐步构造。
3.2 数值稳定的标准 Attention
对第 i 个 query,其最终输出应为全部 N 个 key/value 的数值稳定加权和:
oi=t=1∑Nesit−mit=1∑Nesit−mivt=ℓit=1∑Nesit−mivt
其中:
- 全局最大值 mi=max1≤t≤Nsit,用于数值稳定性;
- 分母 为全局 EXP 求和项,记为 ℓi=∑t=1Nesit−mi;
- 分子各项 为 esit−mivt,即每个 value 向量按指数权重缩放后的结果。
可简写为:
oi=ℓi1t=1∑Nesit−mivt
由于 key/value 按行被切分为多个块(每行若干个 key/value 行向量),无法一次性加载全部 key/value 计算上述求和。因此,oi 必须增量构造:维护一个“当前已处理 key 的累积输出”,每处理一个新块就将其纳入,同时保持与全局计算完全一致的数值稳定性。
如果没看懂这一部分,则参看 解析 FlashAttention(1):从标准 Attention 讲起 详细了解 Attention 计算过程的含义。
3.3 增量更新推导
设当前已处理前 j−1 个 key/value 对(即 k1,…,kj−1 和 v1,…,vj−1),累积状态为 (m,ℓ,oi)。此时 oi 已是这 j−1 个 key 的精确 Attention 输出。现在加入第 j 个 key/value 对 (kj,vj),需要更新 oi 使其覆盖前 j 个 key/value。
3.4 目标形式
覆盖前 j 个 key/value 的正确输出应为:
oi(target)=ℓnew1t=1∑jesit−mnewvt(30)
其中新的全局最大值与全局 EXP 求和项分别为:
mnew=max(m,sij),ℓnew=t=1∑jesit−mnew(31)
式 (30) 的分子为下面两类项的求和:
- 旧项:∑t=1j−1esit−mnewvt(前 j−1 个 key/value 的贡献,需用存量 oi 还原)
- 新项:esij−mnewvj(第 j 个 key/value 的贡献)
3.5 旧项的还原与基准修正
由处理前状态的定义,oi 是前 j−1 个 key/value 在旧全局最大值 m 下的精确输出:
oi=ℓ1t=1∑j−1esit−mvt
其中 ℓ=∑t=1j−1esit−m。
反解未归一化的加权和:
t=1∑j−1esit−mvt=ℓ⋅oi(32)
将式 (32) 的指数基准从 m 修正至新的全局最大值 mnew,两边同乘 em−mnew:
t=1∑j−1esit−mnewvt=ℓ⋅em−mnew⋅oi(33)
式 (33) 即为目标分子中的旧项。其含义为:将已归一化的旧输出还原为未归一化加权和,并将指数基准统一修正至 mnew。
3.6 新项的基准对齐
由于 Bc=1,第 j 个 key/value 块 (kj,vj) 仅含单个 1×d 的 key/value 向量(即仅有一行)。因此该块的局部最大值即为该 key 向量与 query 的内积本身:
mj=sij
局部指数定义为该分数相对于局部最大值的指数:
pij=esij−mj=esij−sij=e0=1
目标分子中的新项为 esij−mnewvj。利用局部统计量将其基准从 mj 对齐至 mnew:
esij−mnewvj=esij−mj⋅emj−mnewvj=pij⋅emj−mnewvj(34)
式 (34) 即为目标分子中的新项。其中 pij 是局部指数(此处恒为 1),emj−mnew 是将局部基准对齐到全局基准的修正因子。
附:Bc=1的情况:
第 j 个 key/value 块 (Kj,Vj) 包含 Bc 个 1×d 的 key/value 向量(即 Bc 行)。该块的局部最大值为块内所有 score 的最大值:
mj=1≤k≤Bcmaxsijk
其中 sijk 表示 query i 与第 j 个块中第 k 个 key 的内积。
局部指数定义为每个 score 相对于该块局部最大值的指数:
pijk=esijk−mj
3.7 分母的更新
根据公式 (31),得覆盖前 j 个 key/value 的全局 EXP 求和项为:
ℓnew=t=1∑jesit−mnew=t=1∑j−1esit−mnew+esij−mnew
上文已得:
ℓ=t=1∑j−1esit−mi
pij=esij−mj
进行对应项的替换,即得新的全局 EXP 求和项:
ℓnew=前 j−1 个 key 的指数和修正ℓ⋅em−mnew+第 j 个 key 的指数和对齐pij⋅emj−mnew(35)
3.8 增量更新式
将式 (33)、(34)、(35) 代入目标形式 (30),得到第 j 轮后 oi 的增量更新:
oi←ℓ⋅em−mnew+pij⋅emj−mnewℓ⋅em−mnewoi+emj−mnewpijvj(36)
随后赋值全局状态:
m←mnew,ℓ←ℓnew
式 (36) 中:
- 分子第一项 ℓ⋅em−mnewoi:前 j−1 个 key 的加权和经基准修正后的结果;
- 分子第二项 emj−mnewpijvj:第 j 个 key 的加权贡献经基准对齐后的结果;
- 分母 ℓnew:前 j 个 key 在统一基准 mnew 下的指数和,作为新的归一化因子。
遍历全部 N 个 key/value 块后,oi 即精确等于第 2 节中定义的标准 Attention 输出。
3.9 从标量更新式到矩阵形式的扩展
当块尺寸 Br,Bc>1 时,Qi∈RBr×d,Kj,Vj∈RBc×d。此时:
Sij=QiKj⊤∈RBr×Bc
由于 softmax 是逐行独立的,Oi 的 Br 行各自维护独立的标量统计量。将式 (36) 的标量运算按行并行打包,即得到矩阵形式。
3.10 局部统计量的向量化
对 Sij 逐行计算:
- 局部行最大值:m~ij=rowmax(Sij)∈RBr
- 局部指数矩阵:P~ij=exp(Sij−m~ij)∈RBr×Bc(逐行广播减法)
- 局部行指数和:ℓ~ij=rowsum(P~ij)∈RBr
3.11 全局统计量的向量化
下标 i 表示第 i 个 query 块 Qi,该块包含 Br 个 query,每行拥有独立的统计量:
minew=max(mi,m~ij)∈RBr(逐元素取最大)
ℓinew=ℓi⊙emi−minew+ℓ~ij⊙em~ij−minew∈RBr(逐元素运算)
其中 ⊙ 为 Hadamard 积。此式即为式 (35) 在 Br 行上的并行版本。
3.12 输出更新的矩阵化与对角矩阵的作用
对标量式 (36) 的分子两项分别作矩阵化:
第一项(旧输出修正):
标量形式为 ℓ⋅em−mnewoi。在 Br>1 时,ℓi 与 emi−minew 均为 Br 维向量,每行 query 拥有独立的标量统计量。为了对 Br 行分别进行缩放而不互相干扰,需要构造对角矩阵:
diag(ℓi)=ℓi,1⋱ℓi,Br∈RBr×Br
左乘对角矩阵的含义:对任意矩阵 X∈RBr×d,左乘 diag(ℓi) 的结果为:
diag(ℓi)X=ℓi,1⋅X1∗⋮ℓi,Br⋅XBr∗
即第 r 行被乘以 ℓi,r,各行之间完全独立。这正是式 (36) 中 ℓ⋅oi 在 Br 行上的并行实现。同理,diag(ℓi)−1 左乘相当于对每一行分别除以 ℓi,r,实现逐行归一化。
因此,旧输出修正项的矩阵形式为:
diag(ℓi)emi−minewOi∈RBr×d
其第 r 行为 ℓi,r⋅emi,r−mi,rnewOi,r∗,与标量形式完全一致。
第二项(新块贡献对齐):
标量形式为 emj−mnewpijvj。在矩阵形式中,新块对 Br 个 query 的未归一化加权和为 P~ijVj∈RBr×d。将其逐行乘以基准对齐因子 em~ij−minew∈RBr(向量与矩阵相乘时逐行广播):
em~ij−minew⊙(P~ijVj)∈RBr×d
其第 r 行为 em~ij,r−mi,rnew⋅(P~ijVj)r∗,对应标量形式的第二项。
合并归一化:
将两项相加后,逐行除以新的全局和 ℓinew。同样使用对角矩阵实现逐行独立归一化:
Oi←diag(ℓinew)−1(diag(ℓi)emi−minewOi+em~ij−minew⊙(P~ijVj))(37)
式 (37) 即为论文 Algorithm 1 中的增量更新公式。遍历所有 key/value 块后,Oi 即为 Br 个 query 对全部 key 的精确全局 Attention 输出,且全程无需将 Sij 或 P~ij 写回 HBM。
4. 映射到 Attention 矩阵形式与 Algorithm 1 逐行详解
上述标量推导直接映射到 FlashAttention 前向伪代码 Algorithm 1。

在 Attention 中,Score 矩阵定义为:
S=QK⊤∈RN×N
其中第 i 行第 j 列元素 Sij=qi⊤kj。softmax 沿行方向进行,即每行独立归一化。因此第 i 行 Si:=[qi⊤k1, qi⊤k2, …, qi⊤kN] 即为上述标量推导中的向量 x。由于不同行之间的 softmax 计算完全独立(无交互),为便于理解,可先考虑 Br=1 的简化情形,即每次只处理一行,再推广到 Br>1 的 Batch 情形。
Algorithm 1 的输入为 Q,K,V∈RN×d 存储在 HBM,SRAM 容量为 M。
第 1 行:设置块大小
Bc=⌈4dM⌉,Br=min(⌈4dM⌉, d)
这里分母取 4d 是因为 SRAM 需要同时容纳 Kj(Bc×d)、Vj(Bc×d)、Qi(Br×d)、Oi(Br×d)以及 Sij(Br×Bc),元素个数总共为 2Bcd+2Brd+BrBc。我们限制:
2Bcd+2Brd+BrBc≤M
然而,实际上代入 Bc 和 Br 的值会发现 2Bcd+2Brd+BrBc 略大于 M,所以论文中设置的 Bc 和 Br 值只是工程上的启发式方法,并不是严格的。
第 2 行:初始化输出与全局统计量
O=0N×d∈RN×d,ℓ=0N∈RN,m=(−∞)N∈RN
三者均存储在 HBM 中。O 是最终输出矩阵,ℓ 是每行的全局 EXP 求和项,m 是每行的全局最大值。初始时 O 为零矩阵,ℓ 为零向量,m 为负无穷向量,表示尚未处理任何分块。
第 3 行:输入矩阵分块
将 Q 沿行方向分为 Tr=⌈N/Br⌉ 块 Q1,…,QTr,每块尺寸 Br×d。将 K 和 V 沿行方向分为 Tc=⌈N/Bc⌉ 块 K1,…,KTc 和 V1,…,VTc,每块尺寸 Bc×d。
第 4 行:输出与统计量分块
将 O 沿行方向分为 Tr 块 O1,…,OTr,每块尺寸 Br×d。将 ℓ 分为 Tr 块 ℓ1,…,ℓTr,每块尺寸 Br×1。将 m 分为 Tr 块 m1,…,mTr,每块尺寸 Br×1。这些分块与 Q 的分块一一对应,便于逐块加载到 SRAM。
下图展示了各个分块的切分和维度:

关键点是:
- K 和 V 的切分相同。
- Q 和 O 的切分相同。
- Q 有多少行,则 l 和 m 的维度就是多少。因为它们对应 Q 每一行的 online softmax 统计量。
如果还是没看懂,则参看 解析 FlashAttention(1):从标准 Attention 讲起 详细了解 Attention 计算过程的含义。
第 5 行:外层循环开始
过程展示:flash_attention_visualization.html
for j=1 to Tc do
外层循环遍历 K 和 V 的分块。每轮迭代处理一个 Kj 和一个 Vj。
第 6 行:加载 Kj,Vj 到 SRAM
将 Kj(Bc×d)和 Vj(Bc×d)从 HBM 加载到 on-chip SRAM。这一步在整个内层循环中只执行一次,意味着同一个 Kj 和 Vj 会被所有 Qi 块复用,显著减少 HBM 读取次数。
第 7 行:内层循环开始
for i=1 to Tr do
内层循环遍历 Q 的分块。每轮迭代处理一个 Qi 块,并更新对应的 Oi,ℓi,mi。
第 8 行:加载 Qi,Oi,ℓi,mi 到 SRAM
将 Qi(Br×d)、Oi(Br×d)、ℓi(Br)、mi(Br)从 HBM 加载到 SRAM。注意 Oi,ℓi,mi 是上一次内层迭代更新后的值;在首次迭代时,它们分别为零矩阵、零向量、负无穷向量。
第 9 行:在 SRAM 中计算局部 score 矩阵
Sij=QiKj⊤∈RBr×Bc
这是矩阵乘法:Qi(Br×d)乘以 Kj⊤(d×Bc),得到 Sij(Br×Bc)。该块仅在 SRAM 中临时存在,绝不写入 HBM,这是 FlashAttention 节省内存的核心。

第 10 行:在 SRAM 中计算局部 softmax 统计量
m~ij=rowmax(Sij)∈RBr,P~ij=exp(Sij−m~ij)∈RBr×Bc,ℓ~ij=rowsum(P~ij)∈RBr
m~ij 是 Sij 每行的最大值,即局部最大值。P~ij 是每行减去该行最大值后的逐元素指数,即局部未归一化指数矩阵。ℓ~ij 是 P~ij 每行的和,即局部 EXP 求和项。
第 11 行:在 SRAM 中更新全局统计量
minew=max(mi,m~ij)∈RBr,ℓinew=emi−minewℓi+em~ij−minewℓ~ij∈RBr
minew 逐元素比较此前全局最大值 mi 与当前分块局部最大值 m~ij,取较大者。ℓinew 将此前全局求和项 ℓi 和当前分块局部求和项 ℓ~ij 分别用指数因子调整到新的全局最大值 minew 基准下,再相加。
第 12 行:在 SRAM 中增量更新输出并写回 HBM
Oi←diag(ℓinew)−1(diag(ℓi)emi−minewOi+em~ij−minewP~ijVj)(30)
推导过程见上文。
第 13 行:将更新后的全局统计量写回 HBM
ℓi←ℓinew,mi←minew
这两个 Br 维向量写回 HBM,供下一次内层迭代或反向传播使用。
第 14 行:end for(内层循环结束)
第 15 行:end for(外层循环结束)
第 16 行:Return O
最终返回的 O 就是精确的 Attention 输出 O=softmax(QK⊤)V。
评论