How#
只要不把 N×N 的 S 和 P 寫進 HBM, 就可以加速
- Tiling —— 分塊 + softmax scaling, 讓 softmax 可以增量計算
- Recomputation —— 只存 O,ℓ,m,反向時在 SRAM 裡重算 S、P
再加上 kernel fusion: 整條 pipeline (matmul → mask → softmax → dropout → matmul) 塞進一個 CUDA kernel
但分塊的情況下就算不出 softmax (需要看整個row), 所以就需要 Online Softmax
Online Softmax#
Naive softmax (2 遍掃描)
yi=∑j=1Vexjexi d_j ← d_{j-1} + e^{x_j} # 第一遍:算分母
y_i ← e^{x_i} / d_V # 第二遍:正規化
但是這種作法是 Numerical-unsafe: exj 在 xj 稍大時直接 overflow (fp32 上 x>88 就爆了), 稍小時 underflow 成 0
Safe softmax (3 遍掃描)
標準修法是減去最大值:
yi=∑j=1Vexj−maxkxkexi−maxkxkfor k ← 1, V: # 第一遍:找 max
d_j ← d_{j-1} + e^{x_j - m_V}
y_i ← e^{x_i - m_V} / d_V
當時主流深度學習框架都在用這個安全版本 (TensorFlow v1.7, PyTorch, MxNet, …)
為了把 mV 找出來所以多了一次掃描, 這個部份沒辦法平行化 (mV 要看過整條向量才知道)
Online softmax (2 遍掃描)
for j ← 1, V: # 第一遍:m 和 d 同時算
d_j ← d_{j-1} × e^{m_{j-1} - m_j} + e^{x_j - m_j} # 合併在這
y_i ← e^{x_i - m_V} / d_V
維護一個不變式, 使用目前為止的最大值對目前的所有元素求和:
dj=k=1∑jexk−mjdj=舊元素,但基準換了k=1∑j−1exk−mj+新元素exj−mj但問題是舊的 dj−1 是以 mj−1 為基準算的,不是 mj, 所以需要 rescale factor: e−mj−1e+mj−1
k=1∑j−1exk−mj=k=1∑j−1exk−mj−1⋅emj−1−mj=dj−1⋅emj−1−mj其中 emj−1−mj 可以提到求和符號外面, 因為它跟 k 無關, 所以可以得:
dj=dj−1⋅emj−1−mj+exj−mj平行化#
Online softmax 還是循序的, 但 GPU 要平行, 所以論文把它推廣成一個二元運算子, 把狀態寫成一對 (m,d),定義:
(midi)⊕(mjdj)=(max(mi,mj)diemi−max(mi,mj)+djemj−max(mi,mj))所以整條向量的統計量就是:
(mVdV)=(x11)⊕(x21)⊕⋯⊕(xV1)運算 ⊕ 是可結合的, 也是可交換的, 可結合 → 可以隨便加括號; 可交換 → 順序無所謂
但 Online softmax 本身還是兩遍, 第一遍算出 (mV,dV), 第二遍才能算 yi=exi−mV/dV, 因為要輸出整條 y 向量, 但每個 yi 都要除以只有掃完才知道的 dV
論文自己也承認這點, 但有指出一個特例:當 softmax 後面接的是 TopK 時, 可以把兩者融合成一遍, 因為 TopK 不需要算出所有 yi
但在 attention 裡, softmax 的輸出 P 從來就不是我們要的東西, 要的是 O=PV; P 只是個中間產物,它馬上就會被 V 加權求和掉, 所以這個跟 TopK 的性質就有點像, 可以塞進online的遞推裡面, 定義一個未正規化的累積輸出:
o~j:=k=1∑jexk−mjvk用一樣的方法可以得到:
o~j=k=1∑j−1exk−mj−1vk⋅emj−1−mj+exj−mjvj=o~j−1⋅emj−1−mj+exj−mjvjo=dVo~V狀態從 (m,d) 擴充成三元組 (m,d,o~),⊕ 也擴充一下:
madao~a⊕mbdbo~b=m:=max(ma,mb)daema−m+dbemb−mo~aema−m+o~bemb−mo~ 那一行跟 d 那一行結構完全相同,所以結合律與交換律自動繼承
Algorithm 1: Forward pass#
分塊大小 (M = SRAM 容量):
Bc=⌈4dM⌉,Br=min(⌈4dM⌉,d)Tr=⌈N/Br⌉ (Q 的塊數), Tc=⌈N/Bc⌉ (K/V 的塊數)
HBM 初始化: O=(0)N×d,ℓ=(0)N,m=(−∞)N
for j = 1 … T_c: ← 外層:K/V
載入 Q_i, O_i, ℓ_i, m_i → SRAM ← O、ℓ、m 存在 HBM,每步都要搬
S_ij = Q_i K_jᵀ ∈ R^{B_r × B_c} (on chip)
P̃_ij = exp(S_ij − m̃_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, FLOPs 為 O(N2d), 額外記憶體 O(N)
FLOPs 的算法: 內層每次做兩個 matmul, 各 O(BrBcd); 內層總共執行 TcTr=⌈N/Bc⌉⌈N/Br⌉ 次, 所以
O(BcBrN2⋅BrBcd)=O(N2d)Algorithm 2:加上 mask 與 dropout 的完整版#
實際的 kernel 還要處理:
S=τQK⊤,Smasked=mask(S),P=softmax(Smasked),Pdropped=dropout(P,p),O=PdroppedV其中 τ 通常是 1/d, Algorithm 1 為了簡潔省略了它
Forward pass 回傳:O,ℓ,m,R。
Algorithm 4: backward pass#
在 SRAM 初始化 dK̃_j = 0, dṼ_j = 0 ← 留在 SRAM 累加
載入 Q_i, O_i, dO_i, dQ_i, ℓ_i, m_i
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
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=P⊤dO⟹dvj=i∑PijdoidP=dOV⊤⟹dPij=doi⊤vjsoftmax 的 Jacobian 是 diag(y)−yy⊤, 所以:
dSi:=(diag(Pi:)−Pi:Pi:⊤)dPi:=Pi:∘dPi:−(Pi:⊤dPi:)Pi:定義一個 Di:=Pi:⊤dPi:, 可以得:
Di=∑jLieqi⊤kjdoi⊤vj=doi⊤∑jLieqi⊤kjvj=doi⊤oi
Di 需要對整列的 P 和 dP 做 reduce, 但這樣根本塞不進 SRAM, 化簡後只需要兩個長度 d 的向量做內積
可以得 dSij=Pij(dPij−Di), 接著可得:
dqi=j∑dSijkj,dkj=i∑dSijqi
I/O Complexity#
設 d≤M≤Nd, 標準 attention 需要 Θ(Nd+N2) 次 HBM 存取, FlashAttention 需要 Θ(N2d2M−1) 次
Proof:
- K 和 V 的每個元素只被載入一次
- Q 和 O 則要被掃過 Tc 遍, 每遍 Θ(Nd) 個元素
- 總計 Θ(Nd+NdTc)=Θ(NdTc)
分塊大小的三個 SRAM 約束:
Kj,VjBcd=O(M),Qi,OiBrd=O(M),SijBrBc=O(M)可以得 Bc=Θ(M/d),Br=Θ(min(M/d, M/Bc))=Θ(min(M/d, d)), 接著可得:
Tc=BcN=Θ(MNd)⟹Θ(NdTc)=Θ(MN2d2)典型情況 d=64–128, M≈100 KB, d2≪M, 所以 HBM 存取少了很多倍 (最多 9×)
下界: 在 M∈[d,Nd] 的範圍內,不存在任何演算法能用 o(N2d2M−1) 次 HBM 存取算出精確 attention
假設存在這樣的演算法, 那取 M=Θ(Nd), 他的存取次數會是 o(N2d2/(Nd))=o(Nd), 但 Q、K、V、O 本身就佔 Nd 的空間且一開始就在 HBM, 任何演算法至少得碰它們一次 → Ω(Nd) (矛盾)
反向的 IO 複雜度同樣是 Θ(N2d2M−1)