上一篇看到,Prefill 有大量 tokens 可以組成大型 GEMM。
但 context 變長後,Attention 會產生另一個問題:
不是算太多,而是中間結果搬太多。
FlashAttention 沒有近似 Attention,也沒有改變模型輸出。它重新安排計算順序,讓 GPU 不必把完整 Attention matrix 寫進 HBM。
Attention 可以拆成三步:
S = QKᵀ / √d
P = Softmax(S)
O = PV
如果 sequence length 是 L,每個 head 的 S 就是 [L, L]。
以 batch 1、32 heads、L=2048、bf16 為例:
S = [1, 32, 2048, 2048]
≈ 0.25 GiB
Eager implementation 會把 S 寫進 HBM,再讀回來做 Softmax;得到 P 後,又要寫回 HBM,再讀回來和 V 相乘。
問題不是公式裡的 FLOPs,而是 S 和 P 這兩個 [L, L] 中間結果。
對一列 scores s,直接計算 exp(s) 可能 overflow,因此通常先減掉整列最大值:
m = max(s)
p = exp(s - m)
l = sum(p)
softmax(s) = p / l
減掉 m 不會改變 Softmax 結果,因為分子和分母同時乘上了相同常數。
但這個寫法看起來必須先看完整列,才能知道 m。
如果我們只載入一小塊 scores,要怎麼知道後面的 block 會不會出現更大的值?
這就是 Online Softmax 要解決的問題。
FlashAttention 固定一個 Q block,接著逐塊讀入 K 和 V。
它不保留先前所有 scores,只為每一列保留三個 running states:
m:目前看過的最大 score
l:目前的 unnormalized weight sum
O:目前的 weighted value sum
初始值是:
m = -∞
l = 0
O = 0
每讀入一組 K_block、V_block,就執行:
S_block = Q_block × K_blockᵀ / √d
m_new = max(m, rowmax(S_block))
α = exp(m - m_new)
P = exp(S_block - m_new)
l = α × l + rowsum(P)
O = α × O + P × V_block
m = m_new
所有 K/V blocks 都處理完後,才做最後 normalization:
Output = O / l
假設前一個 block 的最大值是 m,後來看到更大的 m_new。
先前累積的結果使用舊基準:
exp(score - m)
要把它換到新基準,只要乘上:
α = exp(m - m_new)
因為:
exp(score - m) × exp(m - m_new)
= exp(score - m_new)
所以舊的 l 和 O 都能被 rescale 到新基準,再和目前 block 的結果相加。
這就是 Online Softmax 最重要的一步。
它讓我們不需要保存先前的 scores,也不需要先知道整列的最大值,最後卻仍能得到和標準 Softmax 相同的結果。
FlashAttention 是 exact Attention;Tiling 改變的是執行順序,不是數學結果。
MLC 章節後面介紹的 FA4 又多做了一層 conditional rescaling:基本 Online Softmax 每次遇到更大的 row max 都會更新基準並 rescale O;FA4 會先檢查新舊基準的差距,差距不大時暫時保留舊基準,避免一次 accumulator 的讀取、縮放與寫回。這是建立在相同 recurrence 上的 IO 優化,不影響今天推導的基本結果。
標準做法:
QKᵀ
→ 完整 S 寫入 HBM
→ S 讀回做 Softmax
→ 完整 P 寫入 HBM
→ P 讀回乘 V
FlashAttention:
載入 Q / K / V Tiles
→ On-chip 計算 S_block
→ 更新 m、l、O
→ 丟掉 S_block 和 P_block
→ 下一個 K/V Block
→ 最後只寫回 Output
在任一時間,kernel 只需要保存目前的 tiles 和 row-wise states,不需要 materialize 完整 [L, L] matrix。
同一組測量下,估算的 HBM traffic 是:
| Sequence | Flash | Eager | Scores 大小 | Traffic 倍率 |
|---|---|---|---|---|
| 512 | 5.2 MB | 72.4 MB | 0.02 GiB | 13.8× |
| 2048 | 21.0 MB | 1094.7 MB | 0.25 GiB | 52.2× |
| 4096 | 41.9 MB | 4336.9 MB | 1.00 GiB | 103.4× |
| 8192 | 83.9 MB | 17263.8 MB | 4.00 GiB | 205.8× |
Sequence 越長,避免寫回完整 Scores 的價值就越大。
根據這次提供的 benchmark report,以下數字是在一張閒置的 RTX 6000 Ada 上,以 batch 1、L=2048 執行 attention-only workload 的實測結果;本篇沒有另外重跑這組實驗:
| Backend | Time | Achieved Rate |
|---|---|---|
| Eager Attention | 5.175 ms | 3.3 TFLOP/s |
| PyTorch SDPA(Flash backend) | 0.172 ms | 99.9 TFLOP/s |
兩者計算相同公式、FLOPs 也相同,時間卻相差約 30×。
Eager path 不是算不動,而是反覆 materialize 中間結果;Flash path 把 tiles 和 Online Softmax state 留在 on-chip memory,讓 GPU 把時間花在真正的矩陣運算上。
這個倍率只代表這張 GPU 與這組 shape,但演算法背後的 IO 問題會隨 sequence length 持續放大。
明天會沿著 MLC 的 FlashAttention 章節,先做最基本的 algorithm mapping:
Q 固定成一個 tile;K、V;m、l、O;m_new 與舊 accumulator 的 rescale;l 並寫回 output;先把 block-wise Online Softmax 做對,再談 TMA、Tensor Core、warp specialization 和 FA4 pipeline。
FlashAttention 的核心不是某一條特殊 GPU instruction,而是一個 IO-aware algorithm:
Tiling
+ Online Softmax
+ Running Output Accumulator
= 不需要完整 Attention Matrix
FlashAttention 的重點不是少算,而是少搬。
今天先理解演算法。明天再把這個 recurrence 真正放進 tiled kernel。