解析 FlashAttention(3):FlashAttention-v1 反向传播
前置阅读:解析 FlashAttention(2):FlashAttention-v1 前向传播
FlashAttention-v1 反向传播可视化:flashattention_backward.html
1. 背景 & 动机
1.1 标准反向传播的内存瓶颈
标准 Attention 的前向计算链条为:
S=QK⊤,P=softmax(S),O=PV
其中 Q,K,V∈RN×d,S,P∈RN×N,O∈RN×d。
训练时,损失函数 L 对输出 O 的梯度 dO=∂O∂L∈RN×d 由下游层反向传播而来。为求 dQ,dK,dV∈RN×d,需根据多元函数的链式法则,依次求出 L 对 V,P,S,Q,K 的梯度。
首先给出完整的反向链条,后续第 2 节从 N=2,d=2 的具体例子出发逐步推导:
- dV=P⊤⋅dO,来自 O=PV 对 V 的链式求导;
- dP=dO⋅V⊤,来自 O=PV 对 P 的链式求导;
- dS=P∘(dP−D1⊤),来自 softmax 的 Jacobian,其中 D∈RN,∘ 表示 Hadamard 积(逐元素相乘),1⊤∈R1×N 为全 1 行向量;
- dQ=τ⋅dS⋅K,来自 S=τQK⊤ 对 Q 的求导;
- dK=τ⋅dS⊤⋅Q,来自 S=τQK⊤ 对 K 的求导。
式中的 ⋅ 表示常规矩阵乘法。
核心矛盾:标准实现必须在 HBM 中保存前向的中间矩阵 S 或 P(或两者),供反向传播使用。这导致:
- 内存:额外需要 O(N2) 显存存储 P∈RN×N;
- IO:反向时需要多次从 HBM 读取 P(大小 N×N)和 dP,HBM 访问量同样为 O(N2)。
当序列长度 N 很大时(如 4K、16K、64K),N2 的内存与 IO 开销成为不可承受的瓶颈。
1.2 FlashAttention 反向传播的核心思路
FlashAttention 解决这一问题的思路与前向传播一脉相承——IO 感知 + 重计算(Recomputation)。具体建立在三个观察之上:
观察一:不保存 P,而是重计算 Pij。 前向传播仅保存 O∈RN×d、逐行统计量 (m,ℓ)∈RN、以及随机数种子 R。反向时,将 Qi,Kj 的小块重新加载到 SRAM,利用 (mi,ℓi) 在片上快速恢复出 Pij。
观察二:分块累加梯度。 dKj,dVj 需要累加所有 Qi 带来的贡献。将 Kj,Vj 置于外层循环,其梯度可在 SRAM 中局部累加,内层循环结束后再一次性写回 HBM。
观察三:Softmax 梯度的关键简化。 反向 softmax 通常需要遍历整行 Pi: 计算。FlashAttention 通过代数变形,将这一操作简化为 Di=rowsum(dOi∘Oi),完全避免了对 N 维向量的存储与遍历。
2. 标准 Attention 反向传播推导
本节从 N=2,d=2 的具体例子出发,写出所有矩阵的具体元素,展示链式法则中每个求和符号的来源,最后推广到一般维度。
2.1 符号定义与具体例子设定
设序列长度 N=2,特征维度 d=2。所有矩阵维度如下:
- Q,K,V,O,dO,dQ,dK,dV∈R2×2
- P,S,dS,dP∈R2×2
- D∈R2(列向量)
记 dP=∂P∂L∈RN×N,dS=∂S∂L∈RN×N。
2.2 矩阵乘法 O=PV 的反向传播
设
P=[p11p21p12p22]∈R2×2,V=[v11v21v12v22]∈R2×2
则
O=PV=[p11v11+p12v21p21v11+p22v21p11v12+p12v22p21v12+p22v22]∈R2×2
即
O11O12O21O22=p11v11+p12v21=p11v12+p12v22=p21v11+p22v21=p21v12+p22v22
故
dO=[dO11dO21dO12dO22]=[∂O11∂L∂O21∂L∂O12∂L∂O22∂L]∈R2×2
推导 dV。 元素 v11 仅出现在 O11 与 O21 中。根据链式法则,损失函数 L 对 v11 的梯度通过这两条路径传递:
∂v11∂L=∂O11∂L⋅∂v11∂O11+∂O21∂L⋅∂v11∂O21
由 O11=p11v11+p12v21,得 ∂v11∂O11=p11。由 O21=p21v11+p22v21,得 ∂v11∂O21=p21。代入:
∂v11∂L=dO11⋅p11+dO21⋅p21
同理,对 v12, v21, v22 分别计算梯度:
∂v12∂L=dO12⋅p11+dO22⋅p21
∂v21∂L=dO11⋅p12+dO21⋅p22
∂v22∂L=dO12⋅p12+dO22⋅p22
注意到
P⊤=[p11p12p21p22]∈R2×2
而
dV=[∂v11∂L∂v21∂L∂v12∂L∂v22∂L]
因此得到矩阵形式:
dV=P⊤⋅dO∈R2×2(1)
推导 dP。 元素 p11 仅出现在 O11 与 O12 中:
∂p11∂L=dO11⋅∂p11∂O11+dO12⋅∂p11∂O12=dO11⋅v11+dO12⋅v12
该式为 dO 的第 1 行与 V⊤ 的第 1 列的内积。因此:
dP=dO⋅V⊤∈R2×2(2)
2.3 Softmax 的反向传播(单行情形)
Softmax 的反向是最复杂的一步。首先考虑单行情形,设输入行向量 s=[s1,s2]∈R1×2,softmax 输出行向量 p=[p1,p2]∈R1×2:
p1=exp(s1)+exp(s2)exp(s1),p2=exp(s1)+exp(s2)exp(s2)
已知上游梯度行向量 dp=[dp1,dp2]∈R1×2,待求 ds=[ds1,ds2]∈R1×2。
计算偏导数 ∂sj∂pk。
当 j=k=1 时:
∂s1∂p1=(exp(s1)+exp(s2))2exp(s1)(exp(s1)+exp(s2))−exp(s1)exp(s1)=p1(1−p1)
当 j=2,k=1 时:
∂s2∂p1=(exp(s1)+exp(s2))20⋅(exp(s1)+exp(s2))−exp(s1)exp(s2)=−p1p2
同理:
∂s1∂p2=−p2p1,∂s2∂p2=p2(1−p2)
组装 Jacobian 矩阵 J=∂s∂p。 将所有偏导数排列成矩阵 J∈R2×2,其中第 k 行第 j 列为 ∂sj∂pk:
J=[∂s1∂p1∂s1∂p2∂s2∂p1∂s2∂p2]=[p1(1−p1)−p2p1−p1p2p2(1−p2)]=diag(p)−p⊤p
其中 diag(p)∈R2×2 为以 p 元素为对角元的对角矩阵,p⊤p∈R2×2 为列向量与行向量的外积。
应用链式法则:
∂s∂L=∂p∂L⋅∂s∂p
即
ds=dp⋅J=[dp1,dp2]⋅[p1(1−p1)−p2p1−p1p2p2(1−p2)]
计算第一个分量 ds1:
ds1=dp1⋅p1(1−p1)+dp2⋅(−p2p1)=p1(dp1−p1dp1−p2dp2)
定义标量
D=p1dp1+p2dp2=dp⋅p⊤∈R(3)
则
ds1=p1(dp1−D),ds2=p2(dp2−D)
合并为向量形式:
ds=p∘(dp−D⋅1⊤)∈R1×2(4)
其中 ∘ 表示 Hadamard 积(逐元素相乘),1⊤=[1,1]∈R1×2,标量 D 通过广播机制扩展到每个位置。
推广到矩阵形式。 Attention 的 softmax 是逐行独立的,每行具有独立的 si,pi,Di。对第 i 行,定义
Di=j=1∑NdPij⋅Pij
令 D=[D1,D2]⊤∈R2,则矩阵形式的 softmax 梯度为:
dS=P∘(dP−D1⊤)∈R2×2(5)
其中 D1⊤∈R2×2 为外积,第 i 行第 j 列元素为 Di。
2.4 Di 的关键简化
式 (3) 定义的 Di 看似需要遍历整行 pi(长度 N),但 FlashAttention 利用前向输出 O 做了代数简化。以下在 N=2,d=2 的例子上验证。
由式 (2),dP=dO⋅V⊤。写出元素形式:
dp11=dO11⋅v11+dO12⋅v12
dp12=dO11⋅v21+dO12⋅v22
代入 D1 的定义:
D1=p11⋅dp11+p12⋅dp12=p11(dO11v11+dO12v12)+p12(dO11v21+dO12v22)=dO11(p11v11+p12v21)+dO12(p11v12+p12v22)
由 2.2 节,O11=p11v11+p12v21,O12=p11v12+p12v22。因此:
D1=dO11⋅O11+dO12⋅O12
同理,对第 2 行:
D2=dO21⋅O21+dO22⋅O22
上式表明:计算 Di 无需访问 P,仅需 dO 的第 i 行与 O 的第 i 行做逐元素乘积后求和。
写成矩阵形式:
D=rowsum(dO∘O)∈R2(6)
其中 rowsum(⋅) 表示对矩阵的每一行求和,结果为一个列向量。第 i 个元素为 ∑k=1ddOik⋅Oik。
该简化的意义:原本计算 Di 需要存储并遍历 N 维向量 pi;现在仅需两个长度为 d 的向量逐元素乘积后求和,复杂度为 O(d),且无需访问 P。
2.5 dQ 与 dK 的推导
与 2.2 节 中 dV 的推导类似。
由 S=τQK⊤。继续使用 N=2,d=2 的例子。设
Q=[q11q21q12q22]∈R2×2,K=[k11k21k12k22]∈R2×2
则 K⊤=[k11k12k21k22]∈R2×2,且
S11S12S21S22=τ(q11k11+q12k12)=τ(q11k21+q12k22)=τ(q21k11+q22k12)=τ(q21k21+q22k22)
推导 dQ。 元素 q11 仅出现在 S11 与 S12 中。根据链式法则:
∂q11∂L=∂S11∂L⋅∂q11∂S11+∂S12∂L⋅∂q11∂S12=dS11⋅τk11+dS12⋅τk21
该式为 dS 的第 1 行 [dS11,dS12] 与 K 的第 1 列 [k11,k21]⊤ 的内积。因此:
dQ=τ⋅dS⋅K∈R2×2(7)
推导 dK。 元素 k11 仅出现在 S11 与 S21 中:
∂k11∂L=dS11⋅∂k11∂S11+dS21⋅∂k11∂S21=dS11⋅τq11+dS21⋅τq21
该式为 dS⊤ 的第 1 行 [dS11,dS21] 与 Q 的第 1 列 [q11,q21]⊤ 的内积。因此:
dK=τ⋅dS⊤⋅Q∈R2×2(8)
2.6 推广到一般维度
以上推导在 N=2,d=2 的例子中完全成立。推广到任意 N 和 d:
- dV=P⊤⋅dO∈RN×d:V 的第 j 行通过 P 的第 j 列影响所有 N 个输出,因此对 i 求和。
- dP=dO⋅V⊤∈RN×N:P 的 (i,j) 元素通过 V 的第 j 行影响 d 个输出通道,因此对 k 求和。
- D=rowsum(dO∘O)∈RN:第 i 行的 Di 由 dOi∈R1×d 与 Oi∈R1×d 的逐元素乘积之和得到。
- dS=P∘(dP−D1⊤)∈RN×N:softmax 的逐行 Jacobian 推广。
- dQ=τ⋅dS⋅K∈RN×d:Q 的 (i,k) 元素通过 Kjk 影响所有 N 个 Sij,因此对 j 求和。
- dK=τ⋅dS⊤⋅Q∈RN×d:K 的 (j,k) 元素通过 Qik 影响所有 N 个 Sij,因此对 i 求和。
2.7 标准反向传播的总结
将上述链条串联,标准反向传播的计算流程为:
- dV=P⊤⋅dO∈RN×d
- dP=dO⋅V⊤∈RN×N
- D=rowsum(dO∘O)∈RN
- dS=P∘(dP−D1⊤)∈RN×N
- dQ=τ⋅dS⋅K∈RN×d,dK=τ⋅dS⊤⋅Q∈RN×d
内存瓶颈:步骤 1、2、4 都需要完整的 P∈RN×N。若 N=4096,FP16 下 P 占用约 32MB;若 N=65536,则占用约 8GB,这仅仅是中间矩阵。
3. FlashAttention 反向传播:分块与重计算
FlashAttention 的解决策略是不在 HBM 中保存 P,而是将上述推导链条拆解到小块上,在 SRAM 中重计算所需的局部 Pij。
3.1 分块策略的直觉
观察式 (7):dQ=τdSK。将 K 按行切分为 K1,…,KTc,则:
dQ=τj=1∑TcdS:jKj
其中 dS:j∈RN×Bc 是 dS 的第 j 列块。这意味着 dQ 可逐块累加得到。
同理:
dKj=τdS:j⊤Q∈RBc×d,dVj=i=1∑TrPij⊤dOi∈RBc×d
dKj 和 dVj 仅依赖于第 j 个 key/value 块与所有 query 块的交互。因此:
外层循环遍历 Kj,Vj,在 SRAM 中为 dKj,dVj 维护局部累加器;内层循环遍历 Qi,重计算 Pij,更新 dQi,dKj,dVj。
这与前向传播中 Kj,Vj 放在外层循环的逻辑完全一致——都是为了让某个块的梯度在 SRAM 中做局部累加,减少 HBM 写回次数。
3.2 在 SRAM 中重计算 Pij
前向传播保存了逐行的全局 softmax 统计量 (mi,ℓi)∈RBr。反向时,加载 Qi∈RBr×d,Kj∈RBc×d 到 SRAM,重计算局部 score:
Sij=τQiKj⊤∈RBr×Bc
应用 mask 后,利用前向保存的 (mi,ℓi) 恢复全局概率:
Pij=diag(ℓi)−1exp(Sijmasked−mi)∈RBr×Bc(9)
与前向博客的衔接:前向博客式 (37) 中,mi 和 ℓi 是处理完所有 key 块后的全局统计量。式 (9) 正是利用它们,将局部 score 矩阵 Sij∈RBr×Bc 恢复为全局归一化后的概率矩阵 Pij∈RBr×Bc。这里的 Pij 与前向博客中 P~ij 的区别在于:P~ij 是局部指数(未归一化到全局),而反向重计算的 Pij 已经是全局 softmax 的精确结果。
3.3 Dropout 的重播
若前向应用了 dropout,标准实现需要保存 N×N 的 dropout mask。FlashAttention 改为:
- 前向保存伪随机数生成器状态 R;
- 反向时恢复 R,在 SRAM 中重新生成与前向完全相同的 dropout mask Zij∈RBr×Bc;
- 应用 dropout:Pijdropped=Pij∘Zij∈RBr×Bc。
这样无需保存巨大的 mask 矩阵,额外内存仅为 O(1)。
3.4 dV 的分块累加
由式 (1),dV=P⊤dO∈RN×d。在分块形式下,第 j 个 key/value 块对 dV 的贡献为:
dVj=i=1∑Tr(Pijdropped)⊤dOi∈RBc×d
因此在内层循环中,对当前 Qi 块计算:
dV~j←dV~j+(Pijdropped)⊤dOi∈RBc×d(10)
其中 dV~j∈RBc×d 是 SRAM 中的局部累加器。
3.5 dP 与 dS 的分块计算
由式 (2),dP=dOV⊤∈RN×N。在分块形式下:
dPijdropped=dOiVj⊤∈RBr×Bc(11)
还原 dropout 梯度(因为前向 Pijdropped=Pij∘Zij):
dPij=dPijdropped∘Zij∈RBr×Bc(12)
由式 (6),Di 仅依赖于 dOi∈RBr×d 和 Oi∈RBr×d,与 j 无关:
Di=rowsum(dOi∘Oi)∈RBr(13)
最后由式 (5),softmax 梯度的分块形式为:
dSij=Pij∘(dPij−Di)∈RBr×Bc(14)
其中 Di∈RBr 通过广播逐行相减。
3.6 dQ 与 dK 的分块累加
由式 (7) 和 (8),对当前块 (i,j):
dQi←dQi+τdSijKj∈RBr×d(15)
dK~j←dK~j+τdSij⊤Qi∈RBc×d(16)
累加的原因:
- dQi∈RBr×d:第 i 个 query 块与所有 key 块交互,因此每个内层循环 j 只贡献一部分梯度,必须用 ← 累加;
- dK~j∈RBc×d:第 j 个 key 块与所有 query 块交互,因此在内层循环中持续累加,直到内层循环结束才写回 HBM。
4. Algorithm 4 逐行详解
基于第 2、3 节的推导,以下完整解释论文 Algorithm 4 的每一行。

输入:Q,K,V,O,dO∈RN×d(HBM);ℓ,m∈RN(HBM,前向保存的 softmax 统计量);SRAM 容量 M;softmax 缩放常数 τ;mask 函数;dropout 概率 pdrop;前向保存的伪随机数生成器状态 R。
第 1 行:Set RNG state to R
将伪随机数生成器状态恢复为 R。这一步确保反向传播中重新生成的 dropout mask 与前向传播完全一致,从而无需保存 N×N 的 dropout mask 矩阵。
第 2 行:Set block sizes
Bc=⌈4dM⌉,Br=min(⌈4dM⌉,d)
块大小设置与前向 Algorithm 1 完全一致。SRAM 需要同时容纳 Kj,Vj∈RBc×d,Qi,Oi,dOi,dQi∈RBr×d,以及重计算的 Sij,Pij∈RBr×Bc 等。
第 3 行:输入矩阵分块
将 Q∈RN×d 沿行分为 Tr=⌈N/Br⌉ 块 Q1,…,QTr,每块 Qi∈RBr×d。将 K,V∈RN×d 沿行分为 Tc=⌈N/Bc⌉ 块 K1,…,KTc 和 V1,…,VTc,每块 Kj,Vj∈RBc×d。
第 4 行:输出与梯度分块
将 O∈RN×d 沿行分为 Tr 块 O1,…,OTr,每块 Oi∈RBr×d。将 dO∈RN×d 沿行分为 Tr 块 dO1,…,dOTr,每块 dOi∈RBr×d。将 ℓ∈RN 分为 Tr 块 ℓ1,…,ℓTr,每块 ℓi∈RBr。将 m∈RN 分为 Tr 块 m1,…,mTr,每块 mi∈RBr。
第 5 行:初始化梯度矩阵并分块
dQ=0N×d,dK=0N×d,dV=0N×d
三者均存储在 HBM 中。将 dQ∈RN×d 沿行分为 Tr 块 dQ1,…,dQTr,每块 dQi∈RBr×d。将 dK,dV∈RN×d 沿行分为 Tc 块 dK1,…,dKTc 和 dV1,…,dVTc,每块 dKj,dVj∈RBc×d。
第 6 行:for j = 1 to T_c do
外层循环遍历 K 和 V 的分块。每轮迭代处理一个 Kj∈RBc×d 和一个 Vj∈RBc×d,计算它们对 dK 和 dV 的贡献。
第 7 行:加载 Kj,Vj 到 SRAM
将 Kj∈RBc×d 和 Vj∈RBc×d 从 HBM 加载到 on-chip SRAM。这一步在整个内层循环中只执行一次。
第 8 行:初始化局部梯度块
dK~j=0Bc×d,dV~j=0Bc×d
在 SRAM 中为当前 Kj 和 Vj 对应的梯度累加器分配空间并初始化为零。dK~j,dV~j∈RBc×d。
第 9 行:for i = 1 to T_r do
内层循环遍历 Q 的分块。每轮迭代处理一个 Qi∈RBr×d,重计算对应的局部 Pij∈RBr×Bc,并更新 dQi,dK~j,dV~j。
第 10 行:加载 Qi,Oi,dOi,dQi,ℓi,mi 到 SRAM
将 Qi∈RBr×d、Oi∈RBr×d、dOi∈RBr×d、dQi∈RBr×d、ℓi∈RBr、mi∈RBr 从 HBM 加载到 SRAM。
第 11 行:在 SRAM 中重计算局部 score 矩阵
Sij=τQiKj⊤∈RBr×Bc
在 SRAM 中重新计算 Qi∈RBr×d 与 Kj∈RBc×d 的 score。该块仅在 SRAM 中临时存在,绝不写入 HBM。
第 12 行:在 SRAM 中应用 mask
Sijmasked=mask(Sij)∈RBr×Bc
对 score 矩阵应用 mask(如 causal mask 或 padding mask),将需要屏蔽的位置设为 −∞。
第 13 行:在 SRAM 中重计算概率矩阵 Pij
Pij=diag(ℓi)−1exp(Sijmasked−mi)∈RBr×Bc
对应第 3.2 节的式 (9)。利用前向保存的统计量 (ℓi∈RBr,mi∈RBr) 在 SRAM 中精确恢复出概率矩阵 Pij∈RBr×Bc。其中 mi 通过广播机制逐行相减,diag(ℓi)−1∈RBr×Br 实现逐行归一化。
第 14 行:在 SRAM 中重计算 dropout mask
Zij∈RBr×Bc,Zij,r,c={1−pdrop10with prob. 1−pdropwith prob. pdrop
利用恢复的随机数种子 R,生成与前向完全相同的 dropout mask。
第 15 行:在 SRAM 中应用 dropout
Pijdropped=Pij∘Zij∈RBr×Bc
其中 ∘ 表示 Hadamard 积(逐元素相乘)。这是前向 dropout 操作的精确重播。
第 16 行:在 SRAM 中累加 dVj
dV~j←dV~j+(Pijdropped)⊤dOi∈RBc×d
对应第 3.4 节的式 (10)。由 Oi=∑j′Pij′droppedVj′,因此 Vj 对 Oi 的梯度贡献为 (Pijdropped)⊤dOi。遍历所有 i 块后,即得到完整的 dVj。
第 17 行:在 SRAM 中计算 dPijdropped
dPijdropped=dOiVj⊤∈RBr×Bc
对应第 3.5 节的式 (11)。由 Oi=PijdroppedVj+other blocks,对 Pijdropped 求导得 ∂Pijdropped∂L=dOiVj⊤。
第 18 行:在 SRAM 中还原 dropout 梯度
dPij=dPijdropped∘Zij∈RBr×Bc
对应第 3.5 节的式 (12)。由于前向时 Pijdropped=Pij∘Zij,反向传播需要乘回相同的 mask Zij。注意 Zij 中非零元素为 1−pdrop1,因此这一步同时完成了梯度缩放。
第 19 行:在 SRAM 中计算标量 Di
Di=rowsum(dOi∘Oi)∈RBr
对应第 2.4 节的式 (13) 和第 3.5 节的推导。这是反向 softmax 梯度的核心简化——完全避免了对 N 维向量 Pi: 的存储与遍历,仅需两个长度为 d 的向量逐元素乘积后求和。
第 20 行:在 SRAM 中计算 dSij
dSij=Pij∘(dPij−Di)∈RBr×Bc
对应第 3.5 节的式 (14)。其中 Di∈RBr 通过广播机制逐行相减:对第 r 行,dPij,r,:−Di,r,再逐元素乘以 Pij,r,:。这直接对应式 (5) 的矩阵形式,是 softmax 梯度的分块实现。
第 21 行:在 SRAM 中更新 dQi 并写回 HBM
dQi←dQi+τdSijKj∈RBr×d
对应第 3.6 节的式 (15)。由 Sij=τQiKj⊤,对 Qi 求导得 ∂Qi∂L=τdSijKj。由于 Qi 参与所有 j 块的计算,因此使用累加 ←。计算完成后写回 HBM。
第 22 行:在 SRAM 中更新 dK~j
dK~j←dK~j+τdSij⊤Qi∈RBc×d
对应第 3.6 节的式 (16)。由 Sij=τQiKj⊤,对 Kj 求导得 ∂Kj∂L=τdSij⊤Qi。由于 Kj 参与所有 i 块的计算,因此使用累加。注意 dK~j∈RBc×d 暂存在 SRAM 中,待内层循环结束后再统一写回 HBM。
第 23 行:end for(内层循环结束)
第 24 行:将 dK~j,dV~j 写回 HBM
dKj←dK~j,dVj←dV~j
内层循环结束后,当前 Kj 和 Vj 对应的完整梯度已计算完毕,从 SRAM 写回 HBM。
第 25 行:end for(外层循环结束)
第 26 行:Return dQ, dK, dV
最终返回三个梯度矩阵 dQ,dK,dV∈RN×d。
评论