论文链接:
FlashAttention v2 过程演示可视化脚本:flashattention_v2_visualizer.html
一、符号约定(与论文一致)
N 为序列长度,d 为特征维度。Br 为每个 query 块包含的 query 向量行数,Bc 为每个 key/value 块包含的 key/value 向量行数。Tr=⌈N/Br⌉ 为 query 块数,Tc=⌈N/Bc⌉ 为 key/value 块数。
Q,K,V∈RN×d 为输入矩阵。O∈RN×d 为输出。dO∈RN×d 为损失函数对 O 的梯度。
分块矩阵:
- Qi∈RBr×d:第 i 个 query 块。
- Kj,Vj∈RBc×d:第 j 个 key/value 块。
- Oi,dOi∈RBr×d:第 i 个输出块及其梯度。
- dQi∈RBr×d,dKj,dVj∈RBc×d:梯度块。
具体实例: N=4,d=3,Br=2,Bc=2,则 Tr=2,Tc=2。
Q1=[q1q2]∈R2×3,Q2=[q3q4]∈R2×3
K1=[k1k2]∈R2×3,K2=[k3k4]∈R2×3
V1=[v1v2]∈R2×3,V2=[v3v4]∈R2×3
其中 qn,km,vm∈R1×3 均为行向量。
二、前置原理:Online Softmax 的数学推导
FlashAttention 的核心是 Online Softmax。
2.1 问题设定与目标
考虑单个 query 行向量(省略行下标),其标准 Attention 输出为:
o=∑m=1Nexp(sm)∑m=1Nexp(sm)vm∈R1×d
其中 sm=qkm⊤∈R 为第 m 个 key 行向量的分数,vm∈R1×d 为第 m 个 value 行向量。
由于 HBM 容量限制,无法一次性载入全部 N 个 key/value 行向量。将 N 个 key 行向量按每块 Bc 个分成 Tc 个块,第 j 块包含:
Keyj={k(j−1)Bc+1,…,kjBc},Valuej={v(j−1)Bc+1,…,vjBc}
目标: 顺序处理第 1,2,…,Tc 块,每处理完第 j 块后维护三个统计量,使得处理完所有块后,无需重新从头计算,即可得到与上式完全相同的结果。
2.2 统计量的严格定义
单个 query 行向量处理完前 j 个 key/value 块后,定义以下三个统计量:
定义 1(全局最大值):
m(j)=1≤t≤j1≤c≤Bcmaxs(t−1)Bc+c∈R
其中,1≤t≤j 是遍历前 j 个 key 块(共 Tc 个 key 块),1≤c≤Bc 是遍历单个 key 块内的行向量(每个 key 块内有 Bc 个 key 行向量)。合起来的表示就是遍历前 j 个 key 块内的所有 key 行向量。
m(j) 是前 j 个块中所有 key 行向量对应的分数 sm 的最大值。它是数值稳定性的基准,后续所有指数运算均以此最大值为参考点。
定义 2(全局指数和):
ℓ(j)=t=1∑jc=1∑Bcexp(s(t−1)Bc+c−m(j))∈R
ℓ(j) 是以当前全局最大值 m(j) 为基准,前 j 个块中所有分数的指数和。注意分母中的 m(j) 确保了每一项指数均不超过 1(因为 s(t−1)Bc+c≤m(j)),从而避免数值溢出。
定义 3(未归一化加权和):
o(j)=t=1∑jc=1∑Bcexp(s(t−1)Bc+c−m(j))v(t−1)Bc+c∈R1×d
o(j) 是以当前全局最大值 m(j) 为基准,前 j 个块中所有 value 的加权累加和,权重为平移后的指数。
关键观察: 若上述定义成立,则:
ℓ(j)o(j)=∑t=1j∑c=1Bcexp(s(t−1)Bc+c−m(j))∑t=1j∑c=1Bcexp(s(t−1)Bc+c−m(j))v(t−1)Bc+c
分子分母同乘 exp(m(j)):
ℓ(j)o(j)=∑t=1j∑c=1Bcexp(s(t−1)Bc+c)∑t=1j∑c=1Bcexp(s(t−1)Bc+c)v(t−1)Bc+c
当 j=Tc 时,上式恰好等于标准 Attention 输出 o。因此,只要我们能增量维护 m(j),ℓ(j),o(j),最终相除即得正确结果。
2.3 增量更新公式的推导
仍然是单个 query 行向量。假设已处理前 j−1 个 key/value 块,当前维护的统计量为 m(j−1),ℓ(j−1),o(j−1)。现处理第 j 个 key/value 块,该块包含分数 {s(j−1)Bc+1,…,sjBc} 和 value {v(j−1)Bc+1,…,vjBc}。
Step 1:更新全局最大值
第 j 个 key 块的局部最大值为:
m~=1≤c≤Bcmaxs(j−1)Bc+c∈R
新的全局最大值必须在"旧全局最大值"和"新块局部最大值"之间取最大:
m(j)=max(m(j−1),m~)∈R
这直接来自定义 1:全局最大值是前 j 个 key 块中所有分数的最大值。
Step 2:更新全局指数和
由定义 2,新的全局指数和应为:
ℓ(j)=t=1∑jc=1∑Bcexp(s(t−1)Bc+c−m(j))
将求和拆分为旧块(前 j−1 个 key 块)和新块(第 j 个 key 块):
ℓ(j)=旧块贡献t=1∑j−1c=1∑Bcexp(s(t−1)Bc+c−m(j))+新块贡献c=1∑Bcexp(s(j−1)Bc+c−m(j))
处理旧块贡献:
对任意旧项 s(t−1)Bc+c(其中 t≤j−1),利用指数性质 exp(a−b)=exp(a−c)⋅exp(c−b),令 a=s(t−1)Bc+c,b=m(j),c=m(j−1):
exp(s(t−1)Bc+c−m(j))=exp(s(t−1)Bc+c−m(j−1))⋅exp(m(j−1)−m(j))
注意 exp(m(j−1)−m(j)) 与 t,c 无关,可提出求和符号外。因此旧块贡献为:
exp(m(j−1)−m(j))⋅t=1∑j−1c=1∑Bcexp(s(t−1)Bc+c−m(j−1))=exp(m(j−1)−m(j))⋅ℓ(j−1)
处理新块贡献:
新块各项直接以新基准 m(j) 计算:
c=1∑Bcexp(s(j−1)Bc+c−m(j))
合并:
ℓ(j)=exp(m(j−1)−m(j))⋅ℓ(j−1)+c=1∑Bcexp(s(j−1)Bc+c−m(j))
Step 3:更新未归一化加权和
由定义 3,新的未归一化加权和应为:
o(j)=t=1∑jc=1∑Bcexp(s(t−1)Bc+c−m(j))v(t−1)Bc+c
同样拆分为旧块和新块:
o(j)=旧块贡献t=1∑j−1c=1∑Bcexp(s(t−1)Bc+c−m(j))v(t−1)Bc+c+新块贡献c=1∑Bcexp(s(j−1)Bc+c−m(j))v(j−1)Bc+c
处理旧块贡献:
对任意旧项,利用相同的指数性质:
exp(s(t−1)Bc+c−m(j))v(t−1)Bc+c=exp(m(j−1)−m(j))⋅exp(s(t−1)Bc+c−m(j−1))v(t−1)Bc+c
提出公因子 exp(m(j−1)−m(j)):
exp(m(j−1)−m(j))⋅t=1∑j−1c=1∑Bcexp(s(t−1)Bc+c−m(j−1))v(t−1)Bc+c=exp(m(j−1)−m(j))⋅o(j−1)
处理新块贡献:
直接以新基准 m(j) 计算:
c=1∑Bcexp(s(j−1)Bc+c−m(j))v(j−1)Bc+c
合并:
o(j)=exp(m(j−1)−m(j))⋅o(j−1)+c=1∑Bcexp(s(j−1)Bc+c−m(j))v(j−1)Bc+c
2.4 统一更新公式与正确性证明
引入局部平移指数以简化表达:
p~c(j)=exp(s(j−1)Bc+c−m(j))∈R,c=1,…,Bc
p~c(j) 的含义:第 j 个 key 块中第 c 个 key 行向量的分数,相对于当前全局最大值 m(j) 的指数值。
则更新公式统一写为:
m(j)=max(m(j−1),1≤c≤Bcmaxs(j−1)Bc+c)
ℓ(j)=exp(m(j−1)−m(j))⋅ℓ(j−1)+c=1∑Bcp~c(j)
o(j)=exp(m(j−1)−m(j))⋅o(j−1)+c=1∑Bcp~c(j)v(j−1)Bc+c
正确性证明:
由 2.3 节的推导过程,上述递推严格保持了定义 2 和定义 3 所要求的:
ℓ(j)=t=1∑jc=1∑Bcexp(s(t−1)Bc+c−m(j))
o(j)=t=1∑jc=1∑Bcexp(s(t−1)Bc+c−m(j))v(t−1)Bc+c
因此当 j=Tc 时:
ℓ(Tc)o(Tc)=∑m=1Nexp(sm−m(Tc))∑m=1Nexp(sm−m(Tc))vm=∑m=1Nexp(sm)∑m=1Nexp(sm)vm=o
最后一步分子分母同乘 exp(m(Tc))。证毕。
关于定标因子 exp(m(j−1)−m(j)):
该因子的作用是将旧累积量从旧基准 m(j−1) 转换到新基准 m(j)。
- 当 m(j)=m(j−1)(新块未产生更大分数)时,该因子为 1,旧累积量 ℓ(j−1) 和 o(j−1) 无需调整。
- 当 m(j)>m(j−1)(新块产生更大分数)时,该因子小于 1,旧累积量按比例收缩,以确保所有项均以新的更大基准 m(j) 为参考点。
三、从标量到矩阵块:N=4,d=3,Br=2,Bc=2
第二节展示了单个 query 块(若干个 query 行向量)对多个 key 块的标量运算。在 GPU 上,需将多个标量运算组织成矩阵块,由 Tensor Core 批量执行。以下展示这种组织方式。
目标: 计算第 i=1 个 query 块的输出 O1∈R2×3,即同时计算 q1 和 q2 的 Attention 输出。
这本质上是将第二节的标量运算,对 2 个 query 行和 2 个 key/value 块并行执行。
3.1 初始化
对第 i=1 个 query 块,维护以下变量:
-
O1(0)=02×3∈R2×3:未归一化累积输出矩阵。第 r 行 O1(0)[r,:]∈R1×3 对应第 r 个 query 行向量的未归一化加权和 o(0) 的向量形式。初始为零矩阵,因为尚未处理任何 key/value。
-
m1(0)=[−∞−∞]∈R2:全局行最大值向量。第 r 个元素 m1(0)[r]∈R 对应第 r 个 query 行向量的全局最大值 m(0)。初始为负无穷,表示尚未处理任何 key。
-
ℓ1(0)=[00]∈R2:全局行指数和向量。第 r 个元素 ℓ1(0)[r]∈R 对应第 r 个 query 行向量的全局指数和 ℓ(0)。初始为零,因为尚未累加任何指数。
3.2 第 1 轮循环(j=1,处理 K1,V1∈R2×3)
Step 1:计算局部分数矩阵。
S1(1)=Q1K1⊤=[q1k1⊤q2k1⊤q1k2⊤q2k2⊤]∈R2×2
S1(1)[r,c]=qrkc⊤∈R 是标量内积。S1(1) 的第 r 行包含第 r 个 query 行向量与第 j=1 块中所有 Bc=2 个 key 行向量的分数。
用途: S1(1) 是当前 query 块与当前 key 块的所有两两内积,是 softmax 的输入。
Step 2:更新全局行最大值。
m1(1)=max(m1(0), rowmax(S1(1)))∈R2
其中 rowmax(S1(1))∈R2 对 S1(1) 每行取最大,输出长度为 2 的列向量。由于 m1(0)=−∞,故:
m1(1)=rowmax(S1(1))=[max(S1(1)[1,1],S1(1)[1,2])max(S1(1)[2,1],S1(1)[2,2])]
m1(1)[r]∈R 的含义:第 r 个 query 行向量在处理完第 1 个 key 块后,与所有已处理 key 的分数中的最大值。
用途: m1(1) 用于数值稳定性,后续指数运算将以此最大值为基准进行平移。
Step 3:计算局部平移指数矩阵。
P~1(1)=exp(S1(1)−m1(1))∈R2×2
定义: P~1(1) 称为局部平移指数矩阵。其元素 P~1(1)[r,c]=exp(S1(1)[r,c]−m1(1)[r]) 表示:第 r 个 query 行向量与第 c 个 key 行向量的分数,相对于当前全局最大值 m1(1)[r] 的指数值。
运算说明: 此处减法为逐行广播。S1(1)∈R2×2 的第 r 行减去 m1(1)∈R2 的第 r 个元素,得到平移后的分数,再逐元素取指数。
用途: P~1(1) 是当前 key 块对当前 query 块的未归一化注意力权重。由于减去了行最大值,最大元素值为 exp(0)=1,避免了指数溢出。这些权重将用于加权累加 value。
Step 4:更新全局行指数和。
ℓ1(1)=exp(m1(0)−m1(1))⊙ℓ1(0)+rowsum(P~1(1))∈R2
分解说明:
- rowsum(P~1(1))∈R2 对 P~1(1) 每行求和,输出长度为 2 的列向量。第 r 个元素是第 r 个 query 行向量对当前 key 块中所有 key 行向量的平移指数之和。
- exp(m1(0)−m1(1))∈R2 是逐元素的定标因子。由于 m1(0)=−∞,该项为 0。
- ⊙ 为逐元素乘法。
因此:
ℓ1(1)=rowsum(P~1(1))=[∑c=12exp(S1(1)[1,c]−m1(1)[1])∑c=12exp(S1(1)[2,c]−m1(1)[2])]
ℓ1(1)[r]∈R 的含义:第 r 个 query 行向量在处理完第 1 个 key 块后,以当前全局最大值 m1(1)[r] 为基准,与所有已处理 key 块的分数的指数和。
用途: ℓ1(1) 是分母的局部近似。循环结束后,ℓ1(Tc) 将等于标准 softmax 的分母。
Step 5:更新未归一化输出。
O1(1)=diag(exp(m1(0)−m1(1)))−1O1(0)+P~1(1)V1∈R2×3
根据指数函数的倒数性质,存在
exp(a)−1=exp(a)1=exp(−a)
而对角矩阵的逆,就是对每个对角元取倒数,故
O(j)=等价于 diag(exp(m(j)−m(j−1)))diag(exp(m(j−1)−m(j)))−1O(j−1)+P~(j)Vj
分解说明:
- diag(exp(m1(0)−m1(1)))−1∈R2×2 是以定标因子为对角元的对角矩阵的逆。由于 m1(0)=−∞,该对角矩阵为零矩阵,其逆无意义,但乘以 O1(0)=0 后该项整体为零矩阵。
- P~1(1)∈R2×2,V1∈R2×3,矩阵乘法结果 ∈R2×3。
因此:
O1(1)=P~1(1)V1=[∑c=12P~1(1)[1,c]⋅vc∑c=12P~1(1)[2,c]⋅vc]∈R2×3
O1(1)[r,:]∈R1×3 的含义:第 r 个 query 行向量在处理完第 1 个 key 块后,以当前全局最大值 m1(1)[r] 为基准,对所有已处理 key 的 value 的加权累加和。权重为平移后的指数 P~1(1)[r,c]。
用途: O1(1) 是分子的局部近似。循环结束后,O1(Tc)/ℓ1(Tc) 将等于标准 softmax attention 的输出。
3.3 第 2 轮循环(j=2,处理 K2,V2∈R2×3)
Step 1:计算局部分数矩阵。
S1(2)=Q1K2⊤=[q1k3⊤q2k3⊤q1k4⊤q2k4⊤]∈R2×2
S1(2)[r,c] 是第 r 个 query 行向量与第 j=2 块中第 c 个 key 行向量的标量内积。
Step 2:更新全局行最大值。
m1(2)=max(m1(1), rowmax(S1(2)))∈R2
m1(2)[r]∈R 的含义:第 r 个 query 行向量在处理完前 2 个 key 块后,与所有已处理 key 的分数中的最大值。
此处必须分两种情况,因为 m1(2) 的值决定了旧统计量是否需要重新定标。
情况 A:第 r 行最大值未更新,m1(2)[r]=m1(1)[r]。
Step 3:局部平移指数。
P~1(2)[r,c]=exp(S1(2)[r,c]−m1(2)[r])=exp(S1(2)[r,c]−m1(1)[r])
P~1(2)[r,c] 的含义:第 r 个 query 行向量与第 c 个新 key 行向量的分数,相对于当前全局最大值 m1(2)[r] 的指数值。
Step 4:更新全局行指数和。
ℓ1(2)[r]=exp(m1(1)[r]−m1(2)[r])⋅ℓ1(1)[r]+c=1∑2P~1(2)[r,c]=1⋅ℓ1(1)[r]+c=1∑2P~1(2)[r,c]
来源: 第二节统一更新公式的向量化。由于基准未变(m1(2)[r]=m1(1)[r]),定标因子 exp(m1(1)[r]−m1(2)[r])=1,旧指数和 ℓ1(1)[r] 无需调整。新全局指数和为旧和加上新块的平移指数之和。
Step 5:更新未归一化输出。
O1(2)[r,:]=exp(m1(1)[r]−m1(2)[r])⋅O1(1)[r,:]+c=1∑2P~1(2)[r,c]⋅v2+c=O1(1)[r,:]+c=1∑2P~1(2)[r,c]⋅v2+c
来源: 第二节统一更新公式的向量化。由于基准未变,旧加权和 O1(1)[r,:] 无需调整。新未归一化加权和为旧和加上新块的加权 value 之和。
情况 B:第 r 行最大值更新,m1(2)[r]>m1(1)[r]。
Step 3:局部平移指数。
P~1(2)[r,c]=exp(S1(2)[r,c]−m1(2)[r])
Step 4:更新全局行指数和。
ℓ1(2)[r]=exp(m1(1)[r]−m1(2)[r])⋅ℓ1(1)[r]+c=1∑2P~1(2)[r,c]
来源: 第二节统一更新公式的向量化。由于基准提升,旧指数和 ℓ1(1)[r] 必须乘以定标因子 exp(m1(1)[r]−m1(2)[r])<1 进行收缩,以转换为新基准下的表示,再加上新块的贡献。
Step 5:更新未归一化输出。
O1(2)[r,:]=exp(m1(1)[r]−m1(2)[r])⋅O1(1)[r,:]+c=1∑2P~1(2)[r,c]⋅v2+c
来源: 第二节统一更新公式的向量化。旧加权和 O1(1)[r,:] 是以旧基准 m1(1)[r] 计算的,必须乘以相同定标因子 exp(m1(1)[r]−m1(2)[r]) 收缩后,才表示新基准下的加权和,再加上新块的贡献。
3.4 最终归一化与 logsumexp
循环结束(j=Tc=2),执行统一归一化:
O1=diag(ℓ1(2))−1O1(2)∈R2×3
对第 r 行:
O1[r,:]=ℓ1(2)[r]1O1(2)[r,:]=∑j=12∑c=12exp(S1,rc(j)−m1,r(2))∑j=12∑c=12exp(S1,rc(j)−m1,r(2))v(j−1)2+c
分子为以最终全局最大值 m1,r(2) 为基准的加权和,分母为对应指数和,与标准 softmax attention 完全一致。
保存 logsumexp:
L1=m1(2)+log(ℓ1(2))∈R2
定义: L1 为对数指数和向量。第 r 个元素 L1[r]=m1,r(2)+log(ℓ1,r(2)) 是第 r 个 query 的 logsumexp。
用途: 反向传播时,利用 L1 可恢复全局 softmax 概率,无需分别保存 m1(2) 和 ℓ1(2)。由 exp(S1(j)−L1)=exp(S1(j)−m1(2))/ℓ1(2),恰为全局 softmax 概率。
四、FlashAttention-2 前向传播的一般形式(Algorithm 1)

对第 i 个 query 块,定义:
- Si(j)=QiKj⊤∈RBr×Bc:第 i 个 query 块与第 j 个 key 块的局部分数矩阵。元素 Si(j)[r,c] 是第 r 个 query 与第 c 个 key 的标量内积。
- mi(j)∈RBr:全局行最大值向量。第 r 个元素是第 r 个 query 行向量在处理完前 j 个 key 块后的全局最大值。初始 mi(0)=(−∞)Br。
- ℓi(j)∈RBr:全局行指数和向量。第 r 个元素是第 r 个 query 行向量在处理完前 j 个 key 块后,以 mi(j)[r] 为基准的指数和。初始 ℓi(0)=0Br。
- Oi(j)∈RBr×d:未归一化累积输出矩阵。第 r 行是第 r 个 query 行向量在处理完前 j 个 key 块后,以 mi(j)[r] 为基准的 value 加权累加和。初始 Oi(0)=0Br×d。
对 j=1,…,Tc,依次执行:
-
Si(j)=QiKj⊤∈RBr×Bc
-
mi(j)=max(mi(j−1), rowmax(Si(j)))∈RBr
-
P~i(j)=exp(Si(j)−mi(j))∈RBr×Bc(逐行广播减法)
定义: P~i(j) 为局部平移指数矩阵。元素 P~i(j)[r,c]=exp(Si(j)[r,c]−mi(j)[r]) 表示第 r 个 query 行向量与第 c 个 key 的分数相对于当前全局最大值 mi(j)[r] 的指数值。
用途: P~i(j) 是当前 key 块对当前 query 块的未归一化注意力权重。减去行最大值确保数值稳定性,最大元素值为 1。
-
ℓi(j)=exp(mi(j−1)−mi(j))⊙ℓi(j−1)+rowsum(P~i(j))∈RBr
来源: 第二节统一更新公式的向量化。exp(mi(j−1)−mi(j))∈RBr 是逐元素的定标因子。当某行基准提升时,对应元素小于 1,对该行的旧指数和进行收缩;当基准不变时,对应元素为 1,旧指数和保持不变。
-
Oi(j)=diag(exp(mi(j−1)−mi(j)))−1Oi(j−1)+P~i(j)Vj∈RBr×d
来源: 第二节统一更新公式的向量化。左侧项将旧加权和按相同定标因子收缩,右侧项加入新块的加权 value 贡献。
-
循环结束后,统一归一化:
Oi=diag(ℓi(Tc))−1Oi(Tc)∈RBr×d
用途: Oi(Tc) 是未归一化的加权和,ℓi(Tc) 是全局指数和。逐行相除得到标准 softmax attention 输出。
-
保存 logsumexp:
Li=mi(Tc)+log(ℓi(Tc))∈RBr
用途: Li 用于反向传播时恢复全局 softmax 概率,替代了分别保存 mi(Tc) 和 ℓi(Tc)。
五、反向传播的梯度推导
5.1 Softmax 梯度的完整推导
设 p∈RN 为 softmax 输出,s∈RN 为输入:
pm=∑k=1Nexp(sk)exp(sm),m=1,…,N
对 pm 关于 sn 求偏导:
当 m=n 时:
∂sn∂pn=(∑kexp(sk))2exp(sn)⋅∑kexp(sk)−exp(sn)⋅exp(sn)=pn(1−pn)
当 m=n 时:
∂sn∂pm=(∑kexp(sk))20−exp(sm)⋅exp(sn)=−pmpn
合并得:
∂sn∂pm=pm(δmn−pn)
由链式法则:
(ds)n=∂sn∂L=m=1∑N(dp)m∂pm∂L⋅Jacobian∂sn∂pm
代入得:
(ds)n=m=1∑Npm(δmn−pn)(dp)m=pn(dp)n−pnm=1∑Npm(dp)m
定义标量 Dn=∑m=1Npm(dp)m=p⊤dp∈R,则:
ds=p⊙(dp−Dn⋅1N)∈RN
定义: Dn 为softmax 梯度修正项。它是输出梯度 dp 与概率 p 的内积,表示当前 query 行向量对所有 key 的加权梯度贡献。
用途: Dn 用于修正 softmax 的梯度。由于 softmax 的归一化特性,增大某个分数会通过分母影响其他分数,Dn 正是这一耦合效应的量化。
5.2 Di 的代数简化(Algorithm 2 第 4 行)

对第 n 个 query 行,需计算 Dn=∑m=1NPn,m(dP)n,m∈R。
Step 1:求 (dP)n,m。
由 On,t=∑m=1NPn,mVm,t,对 Pn,m 求偏导得 ∂Pn,m∂On,t=Vm,t。由链式法则:
(dP)n,m=t=1∑d∂On,t∂L∂Pn,m∂On,t=t=1∑d(dO)n,tVm,t
定义: (dP)n,m 为概率梯度。它是损失函数通过输出 On 的所有维度回传到概率 Pn,m 的梯度之和。
Step 2:代入 Dn。
Dn=m=1∑NPn,m(t=1∑d(dO)n,tVm,t)
交换求和顺序:
Dn=t=1∑d(dO)n,t(m=1∑NPn,mVm,t)
Step 3:识别 On,t。
括号内正是 On,t=∑m=1NPn,mVm,t,因此:
Dn=t=1∑d(dO)n,tOn,t
分块表达:
对第 i 个 query 块,dOi,Oi∈RBr×d,则:
Di=rowsum(dOi⊙Oi)∈RBr
定义: Di 为分块 softmax 梯度修正向量。第 r 个元素 Di[r] 对应第 i 块中第 r 个 query 行向量的修正项。
用途: 计算 Di 无需访问 P 的任何元素,仅需 dOi 与 Oi 的逐元素乘积,复杂度 O(Brd),完全在 SRAM 内完成。
5.3 具体实例:N=4,d=3,Br=2,Bc=2 的反向传播
预计算(Algorithm 2 第 4 行):
D=rowsum(dO⊙O)∈R4
定义: D∈R4 为全局修正向量,第 n 个元素 D[n]=∑t=1ddOn,tOn,t。
分块为 D1∈R2(对应 O1,dO1)和 D2∈R2(对应 O2,dO2)。
外层循环 j=1(加载 K1,V1∈R2×3 到 SRAM):
初始化局部累加器:
- dK1=02×3∈R2×3:第 1 个 key 块的梯度累加器。
- dV1=02×3∈R2×3:第 1 个 value 块的梯度累加器。
内层循环 i=1(加载 Q1,O1,dO1∈R2×3,L1∈R2,D1∈R2):
-
重算概率(Algorithm 2 第 11 行):
S1(1)=Q1K1⊤∈R2×2
P1(1)=exp(S1(1)−L1)∈R2×2
定义: P1(1) 为重算概率矩阵。元素 P1(1)[r,c]=exp(S1(1)[r,c]−L1[r])。
来源: 由 L1=m1(Tc)+log(ℓ1(Tc)),有 exp(S1(1)[r,c]−L1[r])=exp(S1(1)[r,c]−m1(Tc)[r])/ℓ1(Tc)[r],恰为全局 softmax 概率。
用途: P1(1) 用于后续梯度计算,替代了存储完整的 N×N 概率矩阵。
-
累加 dV1(Algorithm 2 第 12 行):
dV1←dV1+(P1(1))⊤dO1∈R2×3
来源: 由 dvm=∑nPn,mdon,对块内所有行同时计算。(P1(1))⊤∈R2×2 的转置使得行索引从 query 变为 key,与 dO1∈R2×3 相乘后,得到每个 key 行向量对所有 query 行向量的加权梯度贡献。
用途: dV1 累加第 j=1 个 value 块受到的所有 query 块的梯度影响。
-
计算 dP1(1)(Algorithm 2 第 13 行):
dP1(1)=dO1V1⊤∈R2×2
定义: dP1(1) 为分块概率梯度矩阵。元素 dP1(1)[r,c]=∑t=1d(dO1)r,t(V1)c,t 是第 r 个 query 行向量对第 c 个 key 行向量的概率梯度。
来源: 由 (dP)n,m=∑t(dO)n,tVm,t,对块内所有行同时计算即得矩阵乘法 dO1V1⊤。
-
计算 dS1(1)(Algorithm 2 第 14 行):
dS1(1)=P1(1)⊙(dP1(1)−D112⊤)∈R2×2
定义: dS1(1) 为分块分数梯度矩阵。元素 dS1(1)[r,c] 是损失函数对分数 S1(1)[r,c] 的梯度。
来源: 5.1 节 softmax 梯度公式 ds=p⊙(dp−D⋅1) 的直接分块实现。D112⊤∈R2×2 的每一行均为 D1 的对应元素,实现对 dP1(1) 的逐行广播减法。
用途: dS1(1) 将用于通过链式法则回传到 Q 和 K 的梯度。
-
更新 dQ1(Algorithm 2 第 15 行):
dQ1←dQ1+dS1(1)K1∈R2×3
来源: 由 Sn,m=qnkm⊤,链式法则给出 dqn=∑m(dS)n,mkm。对块内所有行同时计算即得矩阵乘法 dS1(1)K1。
用途: dQ1 累加第 i=1 个 query 块受到的所有 key/value 块的梯度影响。由于多个外层循环可能同时更新 dQ1,v2 中使用 atomic adds。
-
累加 dK1(Algorithm 2 第 16 行):
dK1←dK1+(dS1(1))⊤Q1∈R2×3
来源: 由 dkm=∑n(dS)n,mqn,对块内所有行同时计算。转置 (dS1(1))⊤∈R2×2 使得行索引从 query 变为 key,与 Q1∈R2×3 相乘后,得到每个 key 行向量对所有 query 行向量的梯度贡献。
用途: dK1 累加第 j=1 个 key 块受到的所有 query 块的梯度影响。
内层循环 i=2(加载 Q2,O2,dO2∈R2×3,L2∈R2,D2∈R2):
- S2(1)=Q2K1⊤∈R2×2,P2(1)=exp(S2(1)−L2)∈R2×2。
- dV1←dV1+(P2(1))⊤dO2∈R2×3。
- dP2(1)=dO2V1⊤∈R2×2。
- dS2(1)=P2(1)⊙(dP2(1)−D212⊤)∈R2×2。
- dQ2←dQ2+dS2(1)K1∈R2×3。
- dK1←dK1+(dS2(1))⊤Q2∈R2×3。
内层循环结束,将 dK1∈R2×3 和 dV1∈R2×3 写回 HBM。
外层循环 j=2(加载 K2,V2∈R2×3):
初始化 dK2=02×3∈R2×3,dV2=02×3∈R2×3。
内层循环 i=1:
- S1(2)=Q1K2⊤∈R2×2,P1(2)=exp(S1(2)−L1)∈R2×2。
- dV2←dV2+(P1(2))⊤dO1∈R2×3。
- dP1(2)=dO1V2⊤∈R2×2。
- dS1(2)=P1(2)⊙(dP1(2)−D112⊤)∈R2×2。
- dQ1←dQ1+dS1(2)K2∈R2×3。
- dK2←dK2+(dS1(2))⊤Q1∈R2×3。
内层循环 i=2:
- S2(2)=Q2K2⊤∈R2×2,P2(2)=exp(S2(2)−L2)∈R2×2。
- dV2←dV2+(P2(2))⊤dO2∈R2×3。
- dP2(2)=dO2V2⊤∈R2×2。
- dS2(2)=P2(2)⊙(dP2(2)−D212⊤)∈R2×2。
- dQ2←dQ2+dS2(2)K2∈R2×3。
- dK2←dK2+(dS2(2))⊤Q2∈R2×3。
内层循环结束,将 dK2,dV2∈R2×3 写回 HBM。
5.4 一般形式(Algorithm 2)
对 j=1,…,Tc(外层循环),i=1,…,Tr(内层循环):
-
重算概率(Algorithm 2 第 11 行):
Si(j)=QiKj⊤∈RBr×Bc
Pi(j)=exp(Si(j)−Li)∈RBr×Bc
定义: Pi(j) 为重算概率矩阵。由前向保存的 Li∈RBr 和当前分数 Si(j) 恢复全局 softmax 概率。
来源: exp(Si(j)−Li)=exp(Si(j)−mi(Tc))/ℓi(Tc),恰为全局 softmax 概率。
-
累加 dVj(Algorithm 2 第 12 行):
dVj←dVj+(Pi(j))⊤dOi∈RBc×d
来源: dvm=∑nPn,mdon 的分块矩阵形式。SRAM 内维护局部累加器,内层循环结束后写回 HBM。
-
计算 dPi(j)(Algorithm 2 第 13 行):
dPi(j)=dOiVj⊤∈RBr×Bc
定义: dPi(j) 为分块概率梯度矩阵。
来源: (dP)n,m=∑t(dO)n,tVm,t 的分块矩阵形式。
-
计算 dSi(j)(Algorithm 2 第 14 行):
dSi(j)=Pi(j)⊙(dPi(j)−Di1Bc⊤)∈RBr×Bc
定义: dSi(j) 为分块分数梯度矩阵。
来源: 5.1 节 softmax 梯度公式 ds=p⊙(dp−D⋅1) 的分块形式。Di1Bc⊤∈RBr×Bc 实现逐行广播减法。
-
更新 dQi(Algorithm 2 第 15 行):
dQi←dQi+dSi(j)Kj∈RBr×d
来源: dqn=∑m(dS)n,mkm 的分块矩阵形式。使用 atomic adds 支持序列长度维度的并行化。
-
累加 dKj(Algorithm 2 第 16 行):
dKj←dKj+(dSi(j))⊤Qi∈RBc×d
来源: dkm=∑n(dS)n,mqn 的分块矩阵形式。SRAM 内维护局部累加器,内层循环结束后写回 HBM。
六、FlashAttention(v1)与 FlashAttention-2(v2)的核心区别
前向传播。 论文 Section 2.3.1 描述 FlashAttention(v1)的 online softmax 技巧;论文 Algorithm 1 是 FlashAttention-2(v2)的前向传播。v2 的关键调整:第一,延迟输出归一化至循环结束,维护未归一化输出 Oi(j) 而非每轮都除以 ℓi(j);第二,只保存 logsumexp Li 而非分开保存 mi 和 ℓi。
反向传播。 论文 Algorithm 2 是 FlashAttention-2(v2)的反向传播。v2 使用 Li 代替 (mi,ℓi) 来重算概率,其余分块累加逻辑与 v1 类似,但配合了序列长度维度的并行化。

非 matmul FLOPs。 v1 每轮内层循环都执行完整的输出 rescaling(除以当前 ℓ);v2 将除法延迟到循环结束后,循环内仅保留逐元素指数修正,大幅减少了非 matmul 操作。
并行维度。 v1 仅在 batch 和 heads 维度并行;v2 额外增加序列长度维度的并行化,前向将 query 行块分配到不同 thread block,反向将 key/value 列块分配到不同 thread block,通过 atomic adds 协调 dQ 的更新。
Warp 划分。 v1 采用 Split-K 策略(K,V 切分到不同 warp),需通过 shared memory 通信累加中间结果;v2 改为 Split-Q 策略(Q 切分到不同 warp,K,V 共享),warp 间无需通信,消除了 shared memory 读写瓶颈。

理论峰值利用率。 v1 前向约为 30–50%,反向约为 25–35%;v2 前向可达 50–73%,反向可达 63%,单 A100 在 GPT 训练中可达 225 TFLOPs/s。
评论