桌子太小,
櫃子太遠
貼一篇長文章給 AI,它會頓得比平常久很多。原因不只是「字多要多算」,更在於一位算得飛快的工人,被迫一直在小桌子和遠處的大檔案櫃之間來回搬東西。FlashAttention 沒有讓他算得更快,而是讓他少走幾趟。
故事裡會出現的東西
- 牠片段文章被切開後的一小段(token)。
- 工工人負責算數的人,手很快,就是 GPU 的運算單元。
- 桌工作桌工人手邊的小桌子,取東西一伸手就到,但放不了多少。
- 櫃檔案櫃房間另一頭的大櫃子,什麼都放得下,但每次拿放都要走一趟。
- 表關聯表每個片段對每個片段的分數,一張 N×N 的大表。
- 卡累計卡只有三個數字的小卡片,是故事後半的關鍵。
三千字的文章進門,每個人都要回頭看每個人
先弄清楚:那張大表是怎麼長出來的
先快速複習上一集。句子被切成片段,依序走進房間;每個片段照規則手冊做出三樣東西:胸前的名牌 K(我能提供什麼線索)、桌上的資料夾 V(我真正能給的內容)、手上的便條 Q(我現在想找什麼)。輪到誰發問,就拿便條去和房間裡每個人的名牌比對,算出分數,換成比重,按比重把大家的資料夾抄一部分回來。這個動作叫注意力。
上一集我們只看「牠」一個人發問,比對 4 張名牌,得到一列 4 個分數。可是當你把整篇文章一次貼進來,房間裡不只「牠」,每一個片段都在發問。N 個人各拿便條去比對 N 張名牌,就得到 N 列、每列 N 格,一張 N×N 的關聯表。
拉一下人數,看關聯表長多大
每一列,是一個人拿便條比對所有名牌的結果
畫面上 32 個人已經有 1,024 格了。真實情況是一篇文章動輒好幾千個片段,而模型還有幾十層樓,每層樓好幾組人同時比對。這張表到底會多大?拉一下看看。
一張完整的關聯表要佔多少空間
只算一層樓、一組注意力、一張 N×N 的表,每格 2 bytes。這是依公式算的大小,用來感受數量級。
而且這張表算完之後還沒結束。上一集說過,分數要先換成比重(每一列加起來 100%),再按比重抄資料夾。也就是說,關聯表 S 之後還會生出一張同樣大小的比重表 P,然後才得到每個人真正要帶走的內容 O。有兩張這麼大的表要處理,就輪到工人出場了。
回到電腦裡
一次把整段輸入送進模型的階段叫 Prefill。這時所有片段一起發問,Q 和 K 的兩兩配對形成 N×N 的分數矩陣 S;接著對每一列做 softmax 得到權重矩陣 P;最後 P 乘上 V 得到輸出 O。寫成算式就是 S = QKᵀ/√d、P = softmax(S)、O = PV。S 和 P 是這一次運算的中間結果,用完就丟;KV Cache 存的是每個片段的 K 和 V,數量隨 N 線性增加,兩者不要混在一起。
工人算得很快,慢的是走去櫃子的那段路
桌子太小、櫃子太遠,是 GPU 的真實處境
負責算數的是工人。他的手非常快,但工作環境有兩個限制。第一,他手邊的工作桌很小,只放得下一小塊資料。第二,所有放不下的東西都要收到房間另一頭的大檔案櫃,那裡什麼都放得下,可是每拿一次、放一次,都要走一趟。
現在請他處理 8 個片段的注意力:一張 8×8 的關聯表。最直覺的做法是三個步驟各自做完:先把整張關聯表算出來,再把整張表換成比重,最後拿比重去抄資料夾。可是整張表放不上桌子,於是每一步之間,他都得把表送回櫃子,下一步再搬回來。
看工人怎麼在桌子和櫃子之間來回
大檔案櫃放得多,但遠
工人的工作桌一伸手就到,但小
算完一整輪,工人真正在算數的時間很少,大部分時間花在走路:把 64 格的關聯表放進櫃子、拿回來、換成比重再放回去、再拿回來。中間表來回搬了 4 趟、256 個數值,而這些表最後根本沒有人要,真正要的只有每個人抄回來的那一小份內容 O。
你可能會想:走路能有多慢?以 FlashAttention 論文拿來舉例的一張 A100 顯示卡為數字,工作桌(晶片內的高速工作區)的資料吞吐大約是檔案櫃(顯示記憶體)的十幾倍,但容量只有它的千分之一左右。工人算一格的速度,遠快過把一格搬去櫃子再搬回來。片段一多,搬運的時間就蓋過了計算的時間。
回到電腦裡
工人是 GPU 的運算單元;工作桌是晶片內的 SRAM 與暫存器,速度極快但只有幾十 MB;檔案櫃是 HBM 或 GDDR 顯示記憶體,有幾十 GB,但相對慢得多。兩者都在同一張顯示卡上,這裡說的搬運不是 CPU 和 GPU 之間、也不是硬碟和記憶體之間。像注意力這種「算的量不大、搬的量很大」的工作,瓶頸在記憶體頻寬而不是算力,術語叫 memory-bound。
剛才那種「先算完整張 S、寫回、再讀回做 softmax、再寫回 P、再讀回乘 V」的做法,是把三個步驟各自寫成獨立運算的標準實作方式;每個步驟都很快,慢在步驟之間中間表的搬運。
能不能一小塊算到底,根本不把表放進櫃子?
FlashAttention 的想法,和一張只有三個數字的卡片
工人想到一個辦法:既然桌子放得下一小塊,那就一次只拿一小塊人的名牌和資料夾,算出這一小塊的分數之後,不要收進櫃子,直接在桌上接著做完比重和抄資料,把結果累計到每個發問者自己的小卡片上,然後把這塊丟掉,換下一塊。整張關聯表從頭到尾都沒有完整存在過。
可是這裡有個大問題。換成比重的時候,需要知道這一列全部的分數,因為每個人的比重是「自己的分數佔整列的多少」。只看了前 4 個人,後 4 個人的分數還沒出現,比重怎麼可能算得對?
先想一件簡單的事:算平均不用留下每一筆
要算全班平均分數,你不必把每個人的分數都留在桌上,只要記兩個數:目前的總和和目前的人數。每來一個新的人,總和加上他的分數、人數加一,隨時都可以用總和除以人數得到「到目前為止」的平均。最後一個人來完,得到的就是全班平均,和一次算完全部分毫不差。
注意力的比重比平均複雜一點,因為 softmax 要先把分數取指數。指數很容易爆成天文數字,所以實際計算時會先把每個分數減掉目前看到的最大值再取指數。這就多了一個要記的數字:目前的最大分數。於是工人的累計卡上有三個數:
- 最大值 m到目前為止看到的最高分數,用來當縮放的基準。
- 總和 l每個分數減掉 m 再取指數之後,全部加起來。相當於平均裡的「人數」。
- 加權累計 a每個人的「指數 × 資料夾內容」加起來。相當於平均裡的「總和」。
關鍵的一步是:如果新來的這一塊裡出現了更大的分數,基準 m 就要換掉,而先前用舊基準算出來的 l 和 a 就得整批乘上一個縮放係數,換成新基準之後再把新資料加進去。因為縮放是對整批一起做的,比例關係不會壞掉。全部看完之後,用 a 除以 l,就是正確的注意力輸出。
親手比一次:一次看完全部,和分塊累計,答案一樣嗎?
一個發問者,房間裡有 4 個人。你可以拖動每個人的分數;「內容」是資料夾裡的一個數值。
不管一塊幾個人、分數怎麼拖,最後一塊讀完,a ÷ l 都和一次算完的結果一樣。請特別試「讓最後一個人分數最高」:前三塊算完時,第 1 個人的比重看起來很高,第 4 個人一進來,基準換掉,前面的累計整批縮小,比重重新分配。工人從頭到尾沒有猜過任何比重,他保留的是隨時可以繼續更新的中間量,這和「先算一個固定比例再相加」是完全不同的事。
現在把兩種做法放在一起跑
同樣 8 個人的關聯表,兩位工人同時開工
左邊三個步驟分開做,中間表進出櫃子;右邊每次拿 4×4 的一小塊,在桌上做到底,只更新累計卡。這是流程比較,不是速度比賽,步數不能換算成加速倍數。看不清楚時,用「下一步」自己一格一格推。
做法 A:整張表算完再說
算分數 → 存進櫃子 → 拿回來換比重 → 存進櫃子 → 拿回來抄資料。
檔案櫃裡的中間表
做法 B:FlashAttention,一小塊做到底
拿一小塊名牌和資料夾 → 在桌上算分數、更新累計卡 → 丟掉這塊,換下一塊。
檔案櫃裡的中間表
右邊沒有比左邊少算任何一格:64 格分數一格不少地算了,比重也一格不少地算了。差別只在算完的那一刻它們在哪裡:左邊被完整寫進櫃子又搬回來,右邊在桌上用完就丟。工人的手速一樣,少的是走路。
要注意的是,FlashAttention 不是「唯一會切塊的方法」。矩陣乘法本來就會切塊做,關鍵在於它把算分數、換比重、抄資料三個步驟跨過去接在同一塊裡做完,靠累計卡保證正確,所以中間表不必完整出現。另外,「不存中間表」不等於不用記憶體:名牌和資料夾還是要從櫃子讀進來,抄回來的內容 O 還是要寫出去,累計卡也要一點空間,省下的是那兩張 N×N 大表的來回。
回到電腦裡
把好幾個運算步驟合併在同一段程式裡、中間結果不落地,叫做核心融合(kernel fusion);那三個數字的累計法叫線上 softmax(online softmax),它在數學上和一次算完的 softmax 完全等價,只是浮點數的運算順序不同,結果可能有極小的誤差。分成小塊逐一處理叫 tiling(分塊)。FlashAttention 就是把這三件事組合起來的注意力實作,論文標題裡的 exact 指的是:它沒有為了省事刪掉任何一對本來該算的關聯,不是近似、不是稀疏抽樣、也不是量化。
頁面上的方塊順序、每塊的大小、輸出的時機都是教學安排;真實實作會讓很多組工人平行處理不同的發問者,區塊大小依硬體調整,訓練時還會在反向傳播重算一次分數以省下更多記憶體。
FlashAttention 沒有改變的事
把它放回和 KV Cache 相鄰的位置
故事到這裡,最容易犯的錯是把 FlashAttention 想得太萬能。回到開場那篇三千字的文章,模型處理它其實分成兩段,兩段的關聯表長得不一樣:
第一段:整篇文章一起進門(Prefill)
N 個人同時發問,關聯表是 N×N(套上遮罩約一半)。這是 FlashAttention 最有感的地方:表最大,搬運最多。
第二段:一次接一個字(Decode)
每次只有一個新片段發問,拿 KV Cache 裡的名牌比對,關聯表只有 1×N 一列。這一段的瓶頸主要在讀取越來越長的 KV Cache,需要另一種安排(例如把一列切給很多工人分頭算)。
所以有三件事它沒有做到。它沒有取代 KV Cache,逐字生成時前文的名牌仍然要留、仍然要讀。它沒有讓注意力變成線性:N 個人對 N 個人的配對一對都沒少,長文章的計算量仍然是平方成長,省的是搬運不是配對。它也不保證每種情況都快同樣的倍數:Prefill 和 Decode 的表形狀不同,實際收益隨模型、長度、批次和硬體而異。
最後把幾個常一起出現的名詞放回各自的位置。它們都在讓模型服務更有效率,但針對的不是同一個問題:
KV Cache
解決:重算把每個片段做過的名牌和資料夾留下來,後面的片段不必等它們重做。PagedAttention
解決:空間浪費KV Cache 要佔顯示記憶體。按「頁」分配,不必一開始就為一個對話預留一大塊連續空間,空間才能給更多人用。FlashAttention
解決:中間表搬運這一次注意力運算裡,不把 N×N 的分數表和比重表完整寫進顯示記憶體再讀回。Continuous Batching
解決:等待空位很多人同時使用時,有人提早寫完,下一輪就讓排隊的人補進來,不必等整批都結束。像 vLLM 這類推論引擎,做的就是把上面這些組合起來:管理分頁的 KV Cache、安排批次,並依環境挑選注意力的實作(FlashAttention 是其中一種後端)。它是服務層的軟體,不是另一個模型,也不會讓模型變聰明。
| 故事裡的 | 真正的名稱 | 一句話說明 |
|---|---|---|
| 名牌、資料夾、便條 | K、V、Q | 比對線索、要抄的內容、發問的條件(見前一集) |
| 關聯表 | 分數矩陣 S = QKᵀ/√d | N 個 Q 對 N 個 K 的兩兩分數,本次運算的中間結果 |
| 比重表 | 權重矩陣 P = softmax(S) | 每一列加總 100%,也是中間結果 |
| 每個人抄回的內容 | 注意力輸出 O = PV | 真正要留下、往下傳的東西 |
| 工人 | GPU 運算單元 | 算得很快 |
| 工作桌 | 晶片內 SRAM/暫存器 | 極快、很小 |
| 檔案櫃 | 顯示記憶體(HBM/GDDR) | 很大、相對慢;同一張顯示卡上 |
| 走一趟櫃子 | 記憶體讀寫(IO) | 注意力的真正瓶頸,術語叫 memory-bound |
| 一小塊做到底 | tiling + kernel fusion | 分塊,並把三步驟合在同一個核心裡 |
| 累計卡 m、l、a | 線上 softmax(online softmax) | 只留三個量就能得到精確的 softmax 結果 |
| 整篇一起進門 | Prefill | N×N 的表,FlashAttention 收益最明顯 |
| 一次接一個字 | Decode | 1×N 的一列,瓶頸轉為讀 KV Cache |
故事沒有說完的部分
這一頁用的是縮小的教學設定:8 個片段、4×4 的區塊、單一組注意力、關聯表每格 2 bytes;累計卡的 a 只有一個數字,真實的 a 是一整個向量。工人走幾趟、搬幾個數值,是依流程數出來的,不是量測;真實的收益要看實際硬體。
FlashAttention 後來有第二版、第三版,主要改善的是工人之間怎麼分工和平行,核心想法不變。逐字生成階段有專門針對 1×N 形狀的做法(如 Flash-Decoding),把一列切給多組工人平行處理。這些延伸這裡都沒有展開。
四個問題,確認真的看懂
選了之後會告訴你原因
延伸閱讀
- FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness(原始論文:記憶體層級、分塊、線上 softmax、IO 分析)
- Online normalizer calculation for softmax(累計卡三個量的來源)
- FlashAttention-2(工人之間的分工與平行化改善)
- Dao-AILab / flash-attention(官方實作與硬體條件)
- Attention Is All You Need(Q/K/V 與縮放點積注意力)
- Efficient Memory Management for LLM Serving with PagedAttention(分頁 KV Cache)
- vLLM: Attention Backends(注意力後端的選擇條件)
- PyTorch: Flash-Decoding for long-context inference(逐字生成階段的做法)
- Hugging Face:How caching works(KV Cache 與逐 token 生成)
本頁所有關聯表、比重、累計量與搬運次數皆由瀏覽器依教學設定即時計算,不呼叫任何模型或硬體 API,也不傳送資料。