Slide 22
Slide 22 text
KARAKURI Inc.All rights reserved.
FlashAttentionのタイリングアルゴリズム
FlashAttentionは、Online Softmaxの原理をAttention計算全体
に拡張する。 Q, K, V をブロックに分割し、ブロック単位で
Attentionを計算しながら、出力をインクリメンタルに更新す
る。
入力: Q, K, V(HBM上)
出力: O(HBM上)
Q を T_r 行のブロック Q_1, …, Q_(T_c) に、K, V を T_c 行のブ
ロックに分割する。ブロックサイズは、Q_i, K_j, V_j および中間
結果がすべてSRAMに収まるように選択する。
22
FlashAttention Forward Pass(簡易版)
─────────────────────────────
for each Q-block Qᵢ (i = 1, ..., ⌈T/Bᵣ⌉):
HBMからQᵢをSRAMにロード
初期化: Oᵢ ← 0, mᵢ ← -∞, ℓᵢ ← 0
for each KV-block Kⱼ, Vⱼ (j = 1, ..., ⌈T/Bᶜ⌉):
HBMからKⱼ, VⱼをSRAMにロード
① SRAM上で Sᵢⱼ = Qᵢ Kⱼᵀ / √dⱼ を計算
② (必要に応じてcausal maskを適用)
③ ブロック内統計量を計算:
m
̃ ᵢⱼ = rowmax(Sᵢⱼ)
P̃ᵢⱼ = exp(Sᵢⱼ - m
̃ ᵢⱼ) (ブロック内softmax)
ℓ̃ᵢⱼ = rowsum(P̃ᵢⱼ)
④ グローバル統計量を更新:
mᵢ_new = max(mᵢ, m
̃ ᵢⱼ)
ℓᵢ_new = ℓᵢ · exp(mᵢ - mᵢ_new) + ℓ̃ᵢⱼ · exp(m
̃ ᵢⱼ - mᵢ_new)
⑤ 出力をリスケール + 更新:
Oᵢ ← Oᵢ · (ℓᵢ · exp(mᵢ - mᵢ_new) / ℓᵢ_new)
+ P̃ᵢⱼ · exp(m
̃ ᵢⱼ - mᵢ_new) / ℓᵢ_new · Vⱼ
⑥ mᵢ ← mᵢ_new, ℓᵢ ← ℓᵢ_new
OᵢをHBMに書き出す