2022夏天,斯坦福的一幫人發了篇論文,標題叫《FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness》。
這名字聽著挺無聊對吧,一堆縮寫,一個技術名詞,跟每天幾百篇論文里隨便哪篇長得都一樣。
但我看了之后,覺得這件事挺有意思的。
因為它解決了一個幾乎所有做 Transformer 的人都踩過坑的問題:大模型推理慢,到底慢在哪里。
很多人第一反應是"算不動",覺得算力不夠,換更大的 GPU 就行。
但這篇文章的作者 Tri Dao 一幫人說了句不一樣的話:
不是算不動,是數據搬不動。
他們管這個叫IO-Awareness——在寫注意力算法的時候,不要只盯著 FLOP 數看,要把 GPU 顯存層級之間的讀寫次數也算進去。
聽起來像廢話對吧,但就是這個"廢話",讓 Transformer 的訓練速度提升了 3 倍,內存消耗降低了 20 倍,還讓 Transformer 第一次在 64K 長度的序列上跑出了好結果。
64K。
之前所有方法,要么跑不了,要么跑出來跟隨機猜差不多。
這篇文章,我想把這個論文的核心內容,用人話講清楚。
先看最基本的注意力公式:
O = softmax(QK^T / sqrt(d)) V三個矩陣相乘,再加一個 softmax。看起來很簡單對吧。
但問題在于,Q 和 K 的矩陣乘法會生成一個 N×N 的注意力矩陣,其中 N 是序列長度。
這意味著時間和空間復雜度都是 O(N2)。
序列長度翻倍,計算量翻四倍。
這個平方復雜度是 Transformer 的天然缺陷,從 2017 年那篇"Attention Is All You Need"出來就沒變過。
過去幾年,一堆人想解決這個問題。
方案主要分兩類:
第一類,近似注意力。 用稀疏化、低秩分解、核函數近似等等手段,把注意力矩陣從 N×N 壓縮到接近 N×1。
這類方法理論上把計算復雜度降到了線性或近線性,但實際跑起來,墻鐘時間(wall-clock time)并沒有明顯加速。
為什么?
因為很多方案只關注減少 FLOP,忽略了內存訪問的開銷。
第二類,稀疏注意力。 讓每個 token 只關注有限的其他 token,直接剪掉大量注意力連接。
這種方法確實減少了計算量,但稀疏模式本身也有內存訪問的 overhead,而且效果往往不如 dense attention。
這篇文章的作者認為,現有方案沒效果的根本原因不是算法不行,而是沒有考慮 GPU 內存層級的 IO 特性。
現代 GPU 的內存層級是這樣的:

DRAM(系統內存):容量最大(幾十 GB 到幾百 GB),速度最慢(大約 12.8 GB/s)
HBM(高帶寬內存):GPU 顯存,容量中等(40-80 GB),速度中等(大約 1.5-2.0 TB/s)
SRAM(片上緩存):容量最小(A100 每個 SM 大約 192 KB),速度最快(大約 19 TB/s)
從 HBM 到 SRAM,帶寬差了一個數量級。
但現代 GPU 的計算速度已經超過內存速度了。操作越來越被內存訪問(IO)而不是計算本身瓶頸住。
所以,關鍵問題不是 FLOP 多不多,而是有多少數據在 HBM 和 SRAM 之間來回搬。
這就是"IO-Awareness"的核心思想。
標準的注意力實現,通常是三步:
第一步:計算 QK^T
把 Q 和 K 從 HBM 讀到 SRAM,在芯片上算 QK^T,結果寫回 HBM。
這一步產生了一個 N×N 的注意力分數矩陣 S。
第二步:Softmax
把 S 從 HBM 讀出來,逐行做 softmax,得到 P 矩陣。P 再寫回 HBM。
第三步:PV
把 P 和 V 從 HBM 讀出來,在芯片上算 PV,結果寫回 HBM。
這三步看起來很自然對吧,但每步都在做一件事:把中間結果從 HBM 寫出去,再從 HBM 讀進來。
讓我們算一下 HBM 的訪問總量。
前向傳播:
第一步:讀 Q、K,寫 S → O(N2d + N2) 次 HBM 訪問
第二步:讀 S,寫 P → O(N2) 次 HBM 訪問
第三步:讀 P、V,寫 O → O(N2d + N2) 次 HBM 訪問
前向傳播總共:O(N2d + N2) 次 HBM 訪問,是序列長度的平方級。
反向傳播:
反向傳播需要用到前向計算的 S 和 P 矩陣來計算梯度,所以同樣的,要讀 S、P 寫回 dQ、dK、dV。
反向傳播也大約 O(N2d + N2) 次 HBM 訪問。
整個前向 + 反向傳播,HBM 訪問總量大約是 O(N2d) 次。
但輸入 Q、K、V 本身的總大小只有 O(Nd),輸出 O 也只有 O(Nd)。
數據量是 O(Nd) 的東西,為什么要做 O(N2d) 的 HBM 訪問?
多出來的 O(N2) 次訪問,全是用在那個大得離譜的 N×N 注意力矩陣上。
這個矩陣太大,放不下 SRAM,只能在 HBM 和 SRAM 之間反復搬。
這就是標準實現的根本問題。
FlashAttention 的思路很簡單:用兩個經典技術,避免把 N×N 注意力矩陣寫到 HBM 上。
這兩個技術是:
1. Tiling(分塊計算)
2. Recomputation(重計算)
標準 softmax 需要對整行做歸一化,看起來必須把整行讀進來才能算。
但 math 上有個技巧:softmax 可以分塊計算。
具體來說,如果我把一個向量 x 拆成兩段 x1 和 x2,那么整個向量 x 的 softmax 結果,可以用 x1 和 x2 各自的 softmax 統計量(最大值 m 和歸一化因子 ?)來逐步合并。

公式大概是:

這樣,每次處理一個塊,只需要記錄兩個小值(m 和 ?),就能把結果正確合并起來。
所以 FlashAttention 的做法是:
把 K、V 分成多個塊
每次只把一個塊加載到 SRAM
對 Q 的每個塊,和 K 的這個塊算 QK^T
在 SRAM 里算 softmax,更新 m 和 ?
逐步累積輸出,最后寫回 HBM
關鍵:整個過程中,N×N 注意力矩陣從來沒有完整地出現在 HBM 上。
那反向傳播怎么辦?
反向傳播需要用到前向的 S 和 P 矩陣。標準做法是把前向的 S 和 P 存在 HBM 上,反向時直接讀。
但 FlashAttention 說:不存了,反向時重新算。
它只存前向的輸出 O 和 softmax 的統計量 m、?,這兩個東西很小,O(Nd) 的大小。
反向傳播時,從 HBM 讀 Q、K、V,重新在 SRAM 里算 S 和 P,然后再算梯度。
雖然多算了一些 FLOP,但因為避免了從 HBM 讀 N×N 矩陣的開銷,實際運行時間反而更快。
這在學術上叫selective gradient checkpointing——梯度檢查點的一種選擇。
文章給出了一個嚴格分析。
標準注意力的 HBM 訪問次數是 Θ(N2d + N2)。
FlashAttention 的 HBM 訪問次數是 Θ(N2d / M),其中 M 是 SRAM 的大小。
為什么是 N2d / M?
因為 SRAM 能放下大小為 Θ(M) 的 K、V 塊,每次能處理 Θ(M/d) 個 K 行。
對于 N 行的 Q,需要 N / (M/d) = Nd/M 次掃描。每次掃描加載 O(Nd) 數據,所以總共 O(N2d / M) 次 HBM 訪問。
拿 A100 來算:
d = 64(head 維度)
M ≈ 100KB(每個 SM 的 SRAM)
標準注意力:O(N2 × 64) 次 HBM 訪問
FlashAttention:O(N2 × 64 / 100000) ≈ O(N2 × 0.00064) 次 HBM 訪問
HBM 訪問量減少了大約 100 倍。
雖然實際不可能完全達到理論極限(因為 SRAM 利用率、塊大小選擇等因素),但文章實驗顯示前向傳播減少了約 8 倍,反向傳播減少了約 7 倍,合計約 9 倍的 HBM 訪問量降低。
文章還證明了一個有意思的結論:
對于任何精確注意力算法,在所有可能的 SRAM 大小范圍內,不可能漸近地優于 O(N2d / M) 的 HBM 訪問下界。
換句話說,FlashAttention 在這個意義上是最優的。
FlashAttention 不只是精確注意力,還可以擴展到稀疏注意力。
思路很簡單:如果注意力矩陣是塊稀疏的(比如某些塊全是零),那么在 Tiling 循環中直接跳過這些塊就行。
算法跟 FlashAttention 幾乎一樣,只是加了一個 if 判斷:如果當前塊 M_ij = 0,跳過計算。
文章證明了 Block-Sparse FlashAttention 的 HBM 訪問次數是 Θ(N2d · s / M),其中 s 是非零塊的比例。
s 越小,加速越多。
實驗顯示,在 LRA benchmark 上,Block-Sparse FlashAttention 相對于標準 FlashAttention 有 2.8 倍的加速,同時精度相當。

BERT-large:
在 MLPerf 1.1 上,FlashAttention 比 Nvidia 記錄快了 15%(從 20.0 分鐘降到 17.4 分鐘)。
GPT-2 small:
比 HuggingFace 實現快 3.5 倍(從 9.5 天降到 2.7 天)
比 Megatron-LM 快 2.0 倍(從 4.7 天降到 2.7 天)
GPT-2 medium:
比 HuggingFace 實現快 3.0 倍(從 21.0 天降到 6.9 天)
比 Megatron-LM 快 1.7 倍(從 11.5 天降到 6.9 天)
Long-Range Arena:
平均加速 2.4 倍
FlashAttention 不只是更快,還能訓練出更好的模型。
GPT-2 長上下文:
用 FlashAttention 訓練 GPT-2 small,上下文長度從 1K 提升到 4K,仍然比 Megatron 的 1K 版本快 30%,且 perplexity 低了 0.7。
長文檔分類:
在 MIMIC-III(醫療文本分類)和 ECtHR(法律判決分類)上,增加序列長度帶來顯著提升:
MIMIC-III:16K 序列比 512 序列提升 4.3 分
ECtHR:8K 序列比 512 序列提升 8.5 分
PathFinder 挑戰:
Path-X(16K 序列):FlashAttention 的 Transformer 達到 61.4% 準確率,是第一個在這個任務上超過隨機猜測的 Transformer
Path-256(64K 序列):Block-Sparse FlashAttention 達到 63.1% 準確率
在不同序列長度下:
序列長度 128-512:FlashAttention 比 PyTorch 標準實現快 2-3 倍
序列長度 1024-2048:FlashAttention 比所有近似注意力方法都快
內存占用:FlashAttention 比 PyTorch 標準實現低 20 倍,比 Linformer 低 2 倍
A100:2-4 倍加速
RTX 3090:2.5-4.5 倍加速(HBM 帶寬更低,加速效果更明顯)
T4:加速較少(SRAM 更小,塊大小需要更小)
這篇文章的核心貢獻,可以用一句話概括:
寫注意力算法的時候,要把 GPU 內存層級的讀寫開銷也算進去。
這個"IO-Awareness"的思想聽起來簡單,但在深度學習這個領域里,很少有人真正認真對待過。
大家習慣了看 FLOP 數,看理論復雜度,看 benchmark 上的 accuracy。
但 FLOP 不等于 wall-clock time,不等于內存使用量,不等于實際訓練出來的模型質量。
FlashAttention 用兩個經典技術——Tiling 和 Recomputation——把注意力機制的 HBM 訪問量從 O(N2) 降到了 O(N2 / M),在保持精確計算的同時實現了 3 倍的加速和 20 倍的內存節省。
而且它不只是更快,還讓 Transformer 第一次真正具備了建模 64K 長度上下文的能力。
這就是 IO-Awareness 的力量。
以上,既然看到這里了,如果覺得不錯,隨手點個贊、在看、轉發三連吧,如果想第一時間收到推送,也可以給我個星標?~
謝謝你看我的文章,我們,下次再見。