AI 互動教室 ‹ 個人地端大語言模型實作
下載 .py 單獨開啟實驗場 ↗ 留言回報
KV CACHE · 03

看懂 KV Cache

上一課你調 temperature、top-p,看模型怎麼從一排候選字裡挑一個。 這一課往回追一步:那排候選是怎麼算出來的—— 以及為什麼句子越長,每多寫一個字就越貴。 下面是同一句話生成兩次,左邊沒有快取、右邊有。按下去看看紅色的部分:

沒有 KV CACHE
有 KV CACHE

右邊的實驗場是真的 Python(在你的瀏覽器裡跑,不用安裝任何東西)。 首次載入約需 30–60 秒,正好夠你讀完第 1 節。每一格程式碼都能改、能重跑, 改壞了重新整理就復原——這是你的沙盒,盡量玩。

01 · 兩階段

Prefill 讀題,Decode 寫答案

模型回答問題其實分兩步,而且這兩步的個性完全相反。

Prefill(讀題)把你輸入的整段提示詞一次全部讀進去。 所有 token 的 K、V 在同一次矩陣乘法裡算完,位置之間沒有先後依賴—— 這一步平行度高、跑得快。

Decode(寫答案)一次只生一個字,生完再生下一個。 沒得平行:下一個字要看前面所有字。這裡就是 KV Cache 要解決的戰場。

右邊第一張圖把每一步的兩種工作量畫在一起:藍柱是這一步 要算幾個 token 的 K/V,橘柱是這一步要讀多少格快取。 Decode 的藍柱永遠是 1,橘柱卻一路長高—— 每生一個字,都要把整個快取從記憶體讀一遍,只換來一個 token 的計算量。 算力大量閒著,卡住的是記憶體頻寬。

這個「算力閒著等記憶體」的怪現象先記著。第 6 課的投機解碼會回來收割它。

02 · 三個角色

Q、K、V:一次圖書館找書

每個字都會產生三個向量。用找書比喻最好懂:

Q(Query)查詢 —— 「我想找什麼?」

當下這個新字提出的問題。像你走進圖書館時腦中那句「我要找講快取的書」。

K(Key)索引 —— 「我是講什麼的?」

每個字的標籤,像書背上的書名。拿 Q 跟它比對,越相關的越受關注。

V(Value)內容 —— 「我實際攜帶的資訊」

書的內文。比對完之後,按相關度把 V 加權取出來,就是這一步的輸出。

流程一句話:新字的 Q × 前面所有字的 K(算相關度)→ 按相關度加權取出所有字的 V → 得到下一個字。

看清楚誰是誰的:Q 只有一個,是「當下這個新字」的K 和 V 是「前面每一個字」的,是被查的那張表。 右邊第二張圖是真的算出來的一組權重(加起來剛好是 1)—— 往前拉那個下拉選單,你會看到權重的長度跟著變短:新字只看得到自己和前面的字。

03 · 代價

沒有快取:每寫一個字,前面全部重算一遍

K、V 如果不存起來,會發生什麼事?生第 1 個字,要算前面 4 個字的 K、V; 生第 2 個字,前面變 5 個字,再算一次;生第 3 個字,前面 6 個字,又算一次。 「今天天氣」的 K、V 在第一輪就算過了,結果每一輪都陪跑。

這不是多做一點——總計算量隨長度平方成長。右邊第三張圖把兩條線畫在一起: 有快取是一條幾乎貼著地板的直線,沒快取是一條翹起來的曲線。 預設值(提示 200 字、生成 300 字)之下,沒快取要算 105,350 個 token 的 K/V, 有快取只要 500 個——多做了 211 倍的工

把生成長度再拉長一格,紅線的成長速度就明顯不一樣。那是平方項在動。

04 · 快取

算過的存起來,答案一模一樣

KV Cache 做的事簡單得有點無趣:Prefill 算好的 K、V 放進顯示卡記憶體裡的一張表, Decode 每一步只做四件事——

  1. 只為新字算出它的 Q(和它自己的 K、V)
  2. 拿 Q 去跟快取裡所有舊的 K 比對相關度
  3. 按相關度加權取用快取裡的 V,得到輸出
  4. 把新字的 K、V 也塞進快取,留給下一步用

重要的是:快取不是近似、不是壓縮、不會掉品質。 右邊那格真的跑了兩遍——一遍每步重算整串(合計 12 格), 一遍只算新字其餘查表(合計 5 格)——然後比對兩邊的輸出向量: 完全相同,最大差距在 1e-16 這個量級, 那是浮點數的捨入誤差,不是演算法的差別。

那為什麼只存 K、V,不存 Q?

答案不在 Q 有什麼特別,而在誰會被重複讀。 右邊第四張圖統計整趟生成裡每個 token 被讀取的次數:

  • Q:每個 token 都是 1。問完、答案拿到,這一步就結束了。 之後生成任何新字,都是新字自己重新提問,永遠不會回頭用舊的 Q。
  • K、V:堆得老高。生第 10 個字要查它們、第 100 個字還是要查它們—— 只要對話繼續,就一直被用到。

一句話收尾:快取的意義是「存會被重複讀取的東西」。 這條原則到這裡就講完了,剩下兩節都是它的推論。

05 · 跨請求

Prefix Caching:同一個技術,換一個範圍

到目前為止講的 KV Cache,範圍都在同一次請求之內:Decode 每生一個字, 重用這次對話前面算好的 K、V;請求結束,快取通常就丟掉。

Prefix Caching 就是把它延伸到跨請求、跨使用者:算好的 K、V 留著, 下一個請求如果開頭(前綴)一模一樣,這段的 Prefill 直接跳過。 典型場景是所有人共用同一份系統提示詞。

但它有一個硬條件:只有「從第一個 token 開始、完全相同」的前綴才能重用。 為什麼看起來一樣的字不能搬?右邊第五張圖用兩段話驗證:

  • 「你好,天氣如何」對上「你好,明天天氣如何」:前 3 格的 K 差 0.0%—— 逐位元完全相同,這幾格的 Prefill 可以整段跳過。第 4 格起就衝到 70–125%,一個字都救不回來。
  • 理由 1:位置不同。同一段話原封不動往後移 2 格,K 就差了 23%。 前文一模一樣,只是位置變了——位置編碼在算 K 之前就混進去了。
  • 理由 2:前文不同。換一段不同開頭的前文,再看同一個「天」字,K 差 99%。 模型有很多層,深層的 K、V 是「看過前面所有字」之後才算出來的。

所以「天氣如何」這四個字,在兩個開頭不同的請求裡是兩組不同的 K、V。 看起來一樣,搬過去就會算錯。

這也是為什麼 Prefix Caching 的效果完全由命中率決定:省下的比例就是重複前綴的比例。 9 成前綴相同,Prefill 接近 10 倍;只有 2 成相同,大概 1.25 倍。而且是 token 級精確比對—— 開頭多一個空格、多一行時間戳,後面幾千字的共用前綴就全部作廢。

06 · 帳單

代價:拿記憶體換計算

天下沒有白吃的快取。K、V 要放在顯示卡記憶體裡,而且會隨對話變長一直長大。 每個 token 要吃多少,攤開來算就知道:

每 token bytes = 層數 × KV 頭數 × 每頭維度 × 2(K 和 V 兩份) × 2(fp16 每個數 2 bytes) = 32 × 8 × 128 × 2 × 2 = 131,072 bytes = 128 KiB ≈ 0.13 MB

這是 Llama 3 8B(GQA)的配置。0.13 MB 聽起來很小, 但它是每個 token、而且要乘上同時上線的人數

場景tokens × 並發需要的 KV 快取
手邊隨手測一句2k × 10.26 GB
短文本 RAG8k × 11.05 GB
整篇論文分析32k × 14.19 GB
10 人同時上線8k × 1010.49 GB
整本書128k × 116.78 GB
超長 context1M × 1131.07 GB

假設一張 16 GB 的卡、權重吃掉約 6 GB,剩下約 10 GB 給快取—— 紅色那三列就是放不下的那一邊。而它們一點都不極端: 十個人同時用、丟一整本書進去、或是現在流行的百萬 context,都在這三列裡。

右邊第六張圖是可以拉的估算器:換層數、換 KV 頭數、換上下文長度與人數, 看你的設定落在虛線的哪一邊。放不下的時候怎麼辦?砍上下文、砍並發, 或者把快取搬到別的地方去——那是第 5 課的題目。

07 · 實戰

換你動手

右邊最下面有一格「你的實驗區」。三個挑戰,由易到難(notebook 裡都附了折疊解答):

LEVEL 1

回到 6️⃣ 把「KV 頭數」從 8 拉到 32(等於不共用 KV 頭), 看 1M x1 那根柱子從 131 GB 變成幾 GB。

LEVEL 2

回到 3️⃣ 試兩組設定:提示 2000/生成 100,和提示 100/生成 2000。 兩組的 token 總數差不多,但「省下的倍數」差了十倍以上。哪一種情況快取更划算?為什麼?

LEVEL 3

在實驗區換一組你自己的 A/B 句子,找出「從第幾個字開始 K 就不一樣了」。 試著讓兩句共用越長越好的前綴——然後想想:這件事對你寫 prompt 的順序有什麼啟示?

08 · 驗收

情境測驗

離開前試試看:下面的情境都真的會遇到。每題選一個你認為的最佳做法,選了馬上看得到解釋。

Q1 情境題

你要在一張 16 GB 的卡上跑 Llama 3 8B(權重約 6 GB),服務 10 個同時上線的使用者,每人 8k 上下文。上線前最該先算的是什麼?

解釋:吃掉 VRAM 的不只有權重,還有會隨對話一直長大的 KV 快取。這個場景要 10.49 GB,比權重之後剩下的約 10 GB 還多——上線後第一批使用者塞滿就會爆。A 是最常見的誤判,它只算了固定成本、漏掉了隨長度成長的那一項。C 方向錯了:Decode 階段每步只算一個 token 卻要把整個快取讀一遍,卡住的是記憶體(容量與頻寬),不是算力。D 反而讓事情更糟——batch 越大,同時存在的快取越多。

Q2 情境題

你的客服機器人每個請求都帶同一份 1800 字的規則說明,再接上使用者的問題。你想讓 Prefix Caching 幫你跳過那 1800 字的 Prefill,哪個做法最有效?

解釋:可重用的只有「從第一個 token 起完全相同」的那一段,所以共用的東西必須擺在最前面、而且逐字穩定。A 把每次都不同的內容放在最前面,前綴長度直接歸零。C 更糟:那行時間戳每一次請求都不一樣,後面 1800 字的規則因為位置與前文都被它改變,一格都命中不了。D 的中間段落不是前綴,位置與前文都會隨使用者問題浮動,同樣無法重用。

Q3 情境題

同事提議做一個「片段快取池」:把公司地址、產品規格這些常出現的段落各自算好 K/V 存起來,之後不管它們出現在 prompt 的哪個位置,都直接貼進快取跳過那段 Prefill。這個設計會怎樣?

解釋:這是「非前綴不能重用」的正面版本。右邊第五張圖量過:同一段話只是往後移 2 格,K 就差了 23%;換一段不同的前文再看同一個字,差 99%。貼進去不會報錯,模型只會安靜地算錯——這種 bug 比崩潰難抓。A 正是那個誤解:文字一樣不代表 K/V 一樣。B 連 Q 都存更沒用,Q 只在它自己那一步被用一次。D 講的是容量問題,就算記憶體無限大,這個設計還是錯的。

Q4 錯誤診斷

有人看到右邊第四張圖的統計,提議「那就連 Q 也一起存進快取,反正順手」。結果記憶體用得更多,生成速度卻一點都沒變快。最貼切的解釋是?

整趟生成裡每個 token 被讀取的次數(提示 4 個 + 生成 6 個) token K/V 被讀取 Q 被使用 t0 6 1 t1 6 1 t2 6 1 t3 6 1 t4 6 1 t5 5 1 t6 4 1 t7 3 1 t8 2 1 t9 1 1

解釋:右欄那張圖就是答案——Q 那一欄從頭到尾都是 1。快取只對「會被重複讀取的東西」有意義;K、V 每一步都被查,所以值得存,Q 用一次就丟,存下來只是佔記憶體。要避免這類判斷失誤,養成問一句話的習慣:「這個東西之後還會被誰讀?」C 描述的是變慢,但題目說的是沒變快,症狀不符;A 和 D 都不是真正的原因。

Q5 錯誤診斷

上線前有人估了一下 KV 快取,算出 5.24 GB、看起來很安全,結果第 6 個使用者連進來就 OOM。估算哪裡錯了?

per_token = 32 * 8 * 128 * 2 # 層數 x KV 頭數 x 每頭維度 x fp16 兩個 bytes need_gb = per_token * 8000 * 10 / 1e9 print(f"{need_gb:.2f} GB") # → 5.24 GB # 實際:第 6 個使用者連進來就 OOM

解釋:公式裡的 * 2 只算到「fp16 一個數 2 bytes」,漏掉「K 和 V 是兩份」。補上之後每 token 是 131,072 bytes(128 KiB),8k × 10 人 = 10.49 GB,剛好越過那台機器可用的約 10 GB——OOM 完全在預期之內。記憶體估算的兩個常漏項就是「K、V 兩份」和「乘上並發」。A 換算方向相反,會讓數字更小;C 是常見誤解,除非兩個請求共用前綴、否則每個人的快取是各自獨立的;D 是杜撰,fp16 就是 2 bytes。

Python 環境載入中(首次約 30–60 秒)…讀完第 1 節它就好了