Mobile wallpaper
1813 字
9 分鐘

Flash Attention

2026-03-01

Why Flash Attention#

標準的 Attention:

S=QKRN×N,P=softmax(S),O=PV\mathbf{S} = \mathbf{QK}^\top \in \mathbb{R}^{N\times N},\quad \mathbf{P} = \text{softmax}(\mathbf{S}),\quad \mathbf{O} = \mathbf{PV}

其他的高效 attention 研究都是在減少 FLOPS, 但其實 Attention 是 memory bound (一直從 HBM 去讀寫矩陣), Flash Attention 稍微犧牲 FLOPS 讓 I/O 更快

How#

只要不把 N×NN\times N 的 S 和 P 寫進 HBM, 就可以加速

  1. Tiling —— 分塊 + softmax scaling, 讓 softmax 可以增量計算
  2. Recomputation —— 只存 O,,m\mathbf{O}, \ell, m,反向時在 SRAM 裡重算 S、P

再加上 kernel fusion: 整條 pipeline (matmul → mask → softmax → dropout → matmul) 塞進一個 CUDA kernel

但分塊的情況下就算不出 softmax (需要看整個row), 所以就需要 Online Softmax

Online Softmax#

Naive softmax (2 遍掃描)

yi=exij=1Vexjy_i = \frac{e^{x_i}}{\sum_{j=1}^{V} e^{x_j}}
d₀ ← 0
for j ← 1, V:
d_j ← d_{j-1} + e^{x_j} # 第一遍:算分母
for i ← 1, V:
y_i ← e^{x_i} / d_V # 第二遍:正規化

但是這種作法是 Numerical-unsafe: exje^{x_j}xjx_j 稍大時直接 overflow (fp32 上 x>88x > 88 就爆了), 稍小時 underflow 成 0

Safe softmax (3 遍掃描)

標準修法是減去最大值:

yi=eximaxkxkj=1Vexjmaxkxky_i = \frac{e^{\,x_i - \max_k x_k}}{\sum_{j=1}^{V} e^{\,x_j - \max_k x_k}}
m₀ ← -∞
for k ← 1, V: # 第一遍:找 max
m_k ← max(m_{k-1}, x_k)
d₀ ← 0
for j ← 1, V: # 第二遍:算分母
d_j ← d_{j-1} + e^{x_j - m_V}
for i ← 1, V: # 第三遍:正規化
y_i ← e^{x_i - m_V} / d_V

當時主流深度學習框架都在用這個安全版本 (TensorFlow v1.7, PyTorch, MxNet, …)

為了把 mVm_V 找出來所以多了一次掃描, 這個部份沒辦法平行化 (mVm_V 要看過整條向量才知道)

Online softmax (2 遍掃描)

m₀ ← -∞
d₀ ← 0
for j ← 1, V: # 第一遍:m 和 d 同時算
m_j ← max(m_{j-1}, x_j)
d_j ← d_{j-1} × e^{m_{j-1} - m_j} + e^{x_j - m_j} # 合併在這
for i ← 1, V: # 第二遍:正規化
y_i ← e^{x_i - m_V} / d_V

維護一個不變式, 使用目前為止的最大值對目前的所有元素求和:

dj=k=1jexkmjd_j = \sum_{k=1}^{j} e^{\,x_k - m_j}dj=k=1j1exkmj舊元素,但基準換了+exjmj新元素d_j = \underbrace{\sum_{k=1}^{j-1} e^{\,x_k - m_j}}_{\text{舊元素,但基準換了}} + \underbrace{e^{\,x_j - m_j}}_{\text{新元素}}

但問題是舊的 dj1d_{j-1} 是以 mj1m_{j-1} 為基準算的,不是 mjm_j, 所以需要 rescale factor: emj1e+mj1e^{-m_{j-1}}e^{+m_{j-1}}

k=1j1exkmj=k=1j1exkmj1emj1mj=dj1emj1mj\sum_{k=1}^{j-1} e^{\,x_k - m_j} = \sum_{k=1}^{j-1} e^{\,x_k - m_{j-1}} \cdot e^{\,m_{j-1} - m_j} = d_{j-1}\cdot e^{\,m_{j-1} - m_j}

其中 emj1mje^{m_{j-1} - m_j} 可以提到求和符號外面, 因為它跟 kk 無關, 所以可以得:

dj=dj1emj1mj+exjmjd_j = d_{j-1}\cdot e^{\,m_{j-1} - m_j} + e^{\,x_j - m_j}

平行化#

Online softmax 還是循序的, 但 GPU 要平行, 所以論文把它推廣成一個二元運算子, 把狀態寫成一對 (m,d)(m, d),定義:

(midi)(mjdj)=(max(mi,mj)diemimax(mi,mj)+djemjmax(mi,mj))\begin{pmatrix} m_i \\ d_i \end{pmatrix} \oplus \begin{pmatrix} m_j \\ d_j \end{pmatrix} = \begin{pmatrix} \max(m_i, m_j) \\[4pt] d_i\, e^{\,m_i - \max(m_i,m_j)} + d_j\, e^{\,m_j - \max(m_i,m_j)}\end{pmatrix}

所以整條向量的統計量就是:

(mVdV)=(x11)(x21)(xV1)\begin{pmatrix} m_V \\ d_V\end{pmatrix} = \begin{pmatrix} x_1 \\ 1\end{pmatrix} \oplus \begin{pmatrix} x_2 \\ 1\end{pmatrix} \oplus \cdots \oplus \begin{pmatrix} x_V \\ 1\end{pmatrix}

運算 \oplus 是可結合的, 也是可交換的, 可結合 → 可以隨便加括號; 可交換 → 順序無所謂

但 Online softmax 本身還是兩遍, 第一遍算出 (mV,dV)(m_V, d_V), 第二遍才能算 yi=eximV/dVy_i = e^{x_i - m_V}/d_V, 因為要輸出整條 yy 向量, 但每個 yiy_i 都要除以只有掃完才知道的 dVd_V

論文自己也承認這點, 但有指出一個特例:當 softmax 後面接的是 TopK 時, 可以把兩者融合成一遍, 因為 TopK 不需要算出所有 yiy_i

但在 attention 裡, softmax 的輸出 P\mathbf{P} 從來就不是我們要的東西, 要的是 O=PV\mathbf{O} = \mathbf{PV}; P\mathbf{P} 只是個中間產物,它馬上就會被 V\mathbf{V} 加權求和掉, 所以這個跟 TopK 的性質就有點像, 可以塞進online的遞推裡面, 定義一個未正規化的累積輸出:

o~j:=k=1jexkmjvk\tilde{o}_j := \sum_{k=1}^{j} e^{\,x_k - m_j}\,v_k

用一樣的方法可以得到:

o~j=k=1j1exkmj1vkemj1mj+exjmjvj=o~j1emj1mj+exjmjvj\tilde{o}_j = \sum_{k=1}^{j-1} e^{\,x_k - m_{j-1}}v_k \cdot e^{\,m_{j-1}-m_j} + e^{\,x_j - m_j}v_j = \tilde{o}_{j-1}\cdot e^{\,m_{j-1}-m_j} + e^{\,x_j - m_j}\,v_jo=o~VdVo = \frac{\tilde{o}_V}{d_V}

狀態從 (m,d)(m, d) 擴充成三元組 (m,d,o~)(m, d, \tilde{o})\oplus 也擴充一下:

(madao~a)(mbdbo~b)=(m:=max(ma,mb)daemam+dbembmo~aemam+o~bembm)\begin{pmatrix} m_a \\ d_a \\ \tilde{o}_a\end{pmatrix} \oplus \begin{pmatrix} m_b \\ d_b \\ \tilde{o}_b\end{pmatrix} = \begin{pmatrix} m := \max(m_a, m_b) \\[3pt] d_a e^{m_a - m} + d_b e^{m_b - m} \\[3pt] \tilde{o}_a\, e^{m_a - m} + \tilde{o}_b\, e^{m_b - m}\end{pmatrix}

o~\tilde{o} 那一行跟 dd 那一行結構完全相同,所以結合律與交換律自動繼承

Algorithm 1: Forward pass#

分塊大小 (MM = SRAM 容量):

Bc=M4d,Br=min(M4d,  d)B_c = \left\lceil \frac{M}{4d}\right\rceil,\qquad B_r = \min\left(\left\lceil \frac{M}{4d}\right\rceil,\; d\right)

Tr=N/BrT_r = \lceil N/B_r\rceil (Q 的塊數), Tc=N/BcT_c = \lceil N/B_c\rceil (K/V 的塊數)

HBM 初始化: O=(0)N×d\mathbf{O} = (0)_{N\times d}=(0)N\ell = (0)_Nm=()Nm = (-\infty)_N

for j = 1 … T_c: ← 外層:K/V
載入 K_j, V_j → SRAM
for i = 1 … T_r: ← 內層:Q
載入 Q_i, O_i, ℓ_i, m_i → SRAM ← O、ℓ、m 存在 HBM,每步都要搬
S_ij = Q_i K_jᵀ ∈ R^{B_r × B_c} (on chip)
m̃_ij = rowmax(S_ij)
P̃_ij = exp(S_ij − m̃_ij)
ℓ̃_ij = rowsum(P̃_ij)
m_i^new = max(m_i, m̃_ij)
ℓ_i^new = e^{m_i − m_i^new} ℓ_i + e^{m̃_ij − m_i^new} ℓ̃_ij
寫回 O_i ← diag(ℓ_i^new)⁻¹ ( diag(ℓ_i) e^{m_i − m_i^new} O_i
+ e^{m̃_ij − m_i^new} P̃_ij V_j )
寫回 ℓ_i ← ℓ_i^new,m_i ← m_i^new

Algorithm 1 回傳 O=softmax(QK)V\mathbf{O} = \text{softmax}(\mathbf{QK}^\top)\mathbf{V}, FLOPs 為 O(N2d)O(N^2d), 額外記憶體 O(N)O(N)

FLOPs 的算法: 內層每次做兩個 matmul, 各 O(BrBcd)O(B_rB_cd); 內層總共執行 TcTr=N/BcN/BrT_cT_r = \lceil N/B_c\rceil\lceil N/B_r\rceil 次, 所以

O ⁣(N2BcBrBrBcd)=O(N2d)O\!\left(\frac{N^2}{B_cB_r}\cdot B_rB_cd\right) = O(N^2d)

Algorithm 2:加上 mask 與 dropout 的完整版#

實際的 kernel 還要處理:

S=τQK,Smasked=mask(S),P=softmax(Smasked),Pdropped=dropout(P,p),O=PdroppedV\mathbf{S} = \tau\mathbf{QK}^\top,\quad \mathbf{S}^{\text{masked}} = \text{mask}(\mathbf{S}),\quad \mathbf{P} = \text{softmax}(\mathbf{S}^{\text{masked}}),\quad \mathbf{P}^{\text{dropped}} = \text{dropout}(\mathbf{P}, p),\quad \mathbf{O} = \mathbf{P}^{\text{dropped}}\mathbf{V}

其中 τ\tau 通常是 1/d1/\sqrt{d}, Algorithm 1 為了簡潔省略了它

Forward pass 回傳:O,,m,R\mathbf{O}, \ell, m, R

Algorithm 4: backward pass#

for j = 1 … T_c:
載入 K_j, V_j → SRAM
在 SRAM 初始化 dK̃_j = 0, dṼ_j = 0 ← 留在 SRAM 累加
for i = 1 … T_r:
載入 Q_i, O_i, dO_i, dQ_i, ℓ_i, m_i
S_ij = τ Q_i K_jᵀ
S_ij^masked = mask(S_ij)
P_ij = diag(ℓ_i)⁻¹ exp(S_ij^masked − m_i) ← 重算 P(只用 ℓ 和 m)
Z_ij = 重新生成的 dropout mask ← 用存下的 R
P_ij^dropped = P_ij ∘ Z_ij
dṼ_j += (P_ij^dropped)ᵀ dO_i
dP_ij^dropped = dO_i V_jᵀ
dP_ij = dP_ij^dropped ∘ Z_ij
D_i = rowsum(dO_i ∘ O_i)
dS_ij = P_ij ∘ (dP_ij − D_i)
寫回 dQ_i ← dQ_i + τ dS_ij K_j → HBM ← 唯一要反覆讀寫 HBM 的梯度
dK̃_j += τ dS_ijᵀ Q_i ← 留在 SRAM
寫回 dK_j, dV_j → HBM ← 各只寫一次
dV=PdO    dvj=iPijdoid\mathbf{V} = \mathbf{P}^\top d\mathbf{O} \;\Longrightarrow\; dv_j = \sum_i P_{ij}\,do_idP=dOV    dPij=doivjd\mathbf{P} = d\mathbf{O}\mathbf{V}^\top \;\Longrightarrow\; dP_{ij} = do_i^\top v_j

softmax 的 Jacobian 是 diag(y)yy\text{diag}(y) - yy^\top, 所以:

dSi:=(diag(Pi:)Pi:Pi:)dPi:=Pi:dPi:(Pi:dPi:)Pi:dS_{i:} = \big(\text{diag}(P_{i:}) - P_{i:}P_{i:}^\top\big)dP_{i:} = P_{i:} \circ dP_{i:} - (P_{i:}^\top dP_{i:})P_{i:}

定義一個 Di:=Pi:dPi:D_i := P_{i:}^\top dP_{i:}, 可以得:

Di=jeqikjLidoivj=doijeqikjLivj=doioiD_i = \sum_j \frac{e^{q_i^\top k_j}}{L_i}\,do_i^\top v_j = do_i^\top \sum_j \frac{e^{q_i^\top k_j}}{L_i}v_j = do_i^\top o_i

DiD_i 需要對整列的 P\mathbf{P}dPd\mathbf{P} 做 reduce, 但這樣根本塞不進 SRAM, 化簡後只需要兩個長度 dd 的向量做內積

可以得 dSij=Pij(dPijDi)dS_{ij} = P_{ij}(dP_{ij} - D_i), 接著可得:

dqi=jdSijkj,dkj=idSijqidq_i = \sum_j dS_{ij}k_j,\qquad dk_j = \sum_i dS_{ij}q_i

I/O Complexity#

dMNdd \le M \le Nd, 標準 attention 需要 Θ(Nd+N2)\Theta(Nd + N^2) 次 HBM 存取, FlashAttention 需要 Θ(N2d2M1)\Theta(N^2d^2M^{-1})

Proof:

  • K 和 V 的每個元素只被載入一次
  • Q 和 O 則要被掃過 TcT_c 遍, 每遍 Θ(Nd)\Theta(Nd) 個元素
  • 總計 Θ(Nd+NdTc)=Θ(NdTc)\Theta(Nd + Nd\,T_c) = \Theta(Nd\,T_c)

分塊大小的三個 SRAM 約束:

Bcd=O(M)Kj,Vj,Brd=O(M)Qi,Oi,BrBc=O(M)Sij\underbrace{B_cd = O(M)}_{\mathbf{K}_j, \mathbf{V}_j},\qquad \underbrace{B_rd = O(M)}_{\mathbf{Q}_i, \mathbf{O}_i},\qquad \underbrace{B_rB_c = O(M)}_{\mathbf{S}_{ij}}

可以得 Bc=Θ(M/d)B_c = \Theta(M/d)Br=Θ(min(M/d, M/Bc))=Θ(min(M/d, d))B_r = \Theta\big(\min(M/d,\ M/B_c)\big) = \Theta\big(\min(M/d,\ d)\big), 接著可得:

Tc=NBc=Θ ⁣(NdM)    Θ(NdTc)=Θ ⁣(N2d2M)T_c = \frac{N}{B_c} = \Theta\!\left(\frac{Nd}{M}\right) \;\Longrightarrow\; \Theta(Nd\,T_c) = \Theta\!\left(\frac{N^2d^2}{M}\right)

典型情況 d=64d = 64128128, M100M \approx 100 KB, d2Md^2 \ll M, 所以 HBM 存取少了很多倍 (最多 9×)

下界: 在 M[d,Nd]M \in [d, Nd] 的範圍內,不存在任何演算法能用 o(N2d2M1)o(N^2d^2M^{-1}) 次 HBM 存取算出精確 attention

假設存在這樣的演算法, 那取 M=Θ(Nd)M = \Theta(Nd), 他的存取次數會是 o(N2d2/(Nd))=o(Nd)o(N^2d^2/(Nd)) = o(Nd), 但 Q、K、V、O 本身就佔 NdNd 的空間且一開始就在 HBM, 任何演算法至少得碰它們一次 → Ω(Nd)\Omega(Nd) (矛盾)

反向的 IO 複雜度同樣是 Θ(N2d2M1)\Theta(N^2d^2M^{-1})

Reference#

Flash Attention
https://blog.cyberangel.work/posts/ml-flashattention/
作者
Ethan Lai
發布於
2026-03-01
許可協議
CC BY-NC-SA 4.0

評論區

目錄