iT邦幫忙

2026 iThome 鐵人賽

DAY 12
0
Build on Google AI

在贏家書寫歷史之前:我所看見的 AI,與一位工程師共舞著系列 第 12 篇

Day 12 | 明天,我要和昨天的 Token 約會(上):KV Cache

  • 分享至 

  • xImage
  •  

這篇其實就是在看完 大金老師的 KV Cache 課程 後,整理的筆記與心得。

昨天我們提到,LLM 在推論時總共會分成兩個階段:Prefill 以及 Decode 階段。值得注意的是,模型在訓練時在 Multi-Head Attention 中真正學習到的參數,其實只有 轉換矩陣(W_q, W_k, W_v ... 等等),並不是計算過後的 Q、K、V。也就是說,不論是在 Prefill 還是 Decode 階段,所有的 Q、K、V 向量通通都是透過轉換矩陣運算出來的中間產物,而模型的本質只是把一段 Input 算好後變成 Output。

那問題就來了:

  1. Prefill 資訊保留:Prefill 計算好的這些中間產物,代表的是「模型對於整段上下文的理解」。如果沒有把它存下來,那進入 Decode 階段每猜一個字,前面的 Prompt 就等於要重新算一次,那前面的 Prefill 根本是在算心酸的,等於他也只是一次 Decode 階段而已,根本不需要分什麼 Prefill 和 Decode,所以不可能是這樣。
  2. Decode 重複運算:Decode 階段是透過「不斷預測下一個 Token,並把 Token 接到句尾,再預測下一個」的方式進行推論,所以會一直反覆把幾乎同一段話丟進模型。除了最後一個 Token 是新加進來的以外,前面所有的 Input Tokens 都長得一樣、順序也一樣,相當於一直在做極度重複的運算,而且重複幾百千萬回。

這兩個問題帶來的計算開銷非常龐大。為了解決這個問題,最直覺的解法就是:把之前算過的中間結果存下來,給未來 Decode 時重複利用,這套機制就是「KV Cache」!

那看名字也知道,KV 就是 Key 跟 Value 這兩個向量,接下來的問題就是 「為什麼只有 KV 被 Cache?」 回到 Day09 跟 Day10 的架構圖來看的話其實不難發現,每次我們在「預測下一個字」的時候,我們都只透過當下 最後一個 Token 的 Q 以及 前面所有 Tokens 的 KV 來進行運算。這就是為什麼我們只存 KV,而 Q 可以用完即丟(因為後面就真的不需要哈哈~)


一、1 個 Token 佔多少 Byte?

以 Google 開源模型 Gemma-4-31B 為例

在 Hugging Face 上,我們可以直接打開 Google 開源的 Gemma-4-31B-it 的 config.json 設定檔,裡面就會記錄這個模型訓練時帶入的一些參數。我們先假設以標準的 Multi-Head Attention 來計算,也就是假設每個 Q Head 都有專屬的 K 和 V,而我們需要的參數就會是這些:

{
  "num_hidden_layers": 60,            // Transformer 層數
  "num_attention_heads": 32,          // Q 的 Head 數 
  "num_key_value_heads": 32,          // KV 的 Head 數,原本是 16,但由於我們假設 MHA,所以我們讓每個 Q 都有專屬的 K 和 V,所以 Head 數是跟 Q 一樣
  "head_dim": 256,                    // 每個 Head 的維度
  "dtype": "bfloat16",                // 浮點數精度(每個浮點數佔 2 Bytes)
  "max_position_embeddings": 262144   // 官方支援上下文上限(約 256k Tokens)
}

(是不是看到一堆熟悉的名稱!)

因此,在標準 Multi-Head Attention 情況下,1 個 Token 在 Gemma-4-31B 裡要佔多少顯存的容量會是:

https://ithelp.ithome.com.tw/upload/images/20260925/20183607V9RtFuGZ0q.png

在採取 MHA 策略的情況下,單單 1 個 Token,在 GPU 顯存裡就要佔用接近 2 MB!

套回真實對話的情境,一個 Agent 一輪對話的 Context,在顯存所佔用的空間大小就會如下:

對話上下文長度 (Tokens) 情境描述 (Gemini 說的) 累積的 KV Cache 顯存大小 相當於多少顯卡?(Gemini 說的)
1,000 (1k) 一般問答或聊天 ~= 2 GB 一張家用顯卡可能撐得住
8,000 (8k) 閱讀一篇中長篇的文章 ~= 16 GB 一張 16GB 的顯卡直接被灌滿!
32,000 (32k) 一整份專案規格書或程式碼 ~= 64 GB 光是快取就快超過 Gemma-4-31B 模型權重本體!
128,000 (128k) 現代大模型標配的超長對話 ~= 256 GB 塞爆 3 張 80GB 的 A100/H100 都還不夠放!
256,000 (256k) Gemma-4-31B 官方上限 ~= 512 GB 單一使用者的 Agent 就足以讓伺服器 OOM(下班~

在傳統的 Multi-Head Attention 架構下,當你丟給模型非常多的上下文(256k)時,「單一個使用者」的對話快取,就足以吃掉超過 500 GB 顯存!所以實務上絕對不是這樣做,模型開發商肯定不會坐視不管,接下來的幾種策略可以有效降低 KV Cache 所需佔用的空間。


二、降低 KV Cache 空間的常見策略(一):從 Head 數調整

1. Multi-Query Attention (MQA)

首先第一個想法會是「真的需要這麼多 KV 嗎?」、「能不能單一 Token 每一個 Head 都共用同一組 KV?」,既然每個 Q 都配一組專屬的 K 和 V 覺得太吃顯存,那不如讓所有 Q Head 共用同一組 K 和 V

為什麼不是所有 KV 共用同組 Q?
因為要被存下來的是 KV,如果是所有 KV 共用同一組 Q 沒有意義,不會節省任何空間。

  • 做法:保留原本的 32 個 Q Head(多方面的查詢),但把 KV Head 砍到只留 1 組(num_key_value_heads: 1)。
  • 效果:原本計算公式裡的 Head 數直接除以 32,1 個 Token 所需佔用的空間會從 2 MB 降到約 62.5 KB。原本 256k Tokens 的上下文,現在只要約 16 GB。
  • 缺點:模型在理解複雜邏輯和多重語意時的能力容易打折,畢竟做法有點太極端了,在 Gemma 中 1 組 KV 不過就是兩個 256 維的向量,很難把複雜語意表示清楚。

2. Grouped-Query Attention (GQA)

既然 MHA(32:32)顯存太肥,MQA(32:1)又太激進,那折衷方案就是「分組共用」,也就是目前最常用的 Grouped-Query Attention (GQA)。

  • 做法:把 Q Head 分組,每一組共用一對 K 和 V。
  • 以 Google Gemma-4-31B 的真實設定為例:
    "num_attention_heads": 32,
    "num_key_value_heads": 16,
    
    沒錯,Gemma-4-31B 原本就是用 GQA!它把 32 個 Q Head 分成 16 組,每 2 個 Q Head 共用一組 KV(2:1)。 換算下來 1 個 Token 需要的空間就變成原本的一半 1MB。GQA 是個能在不掉智商的前提下,大幅省下所需佔用顯存空間的一種策略,現在主流模型基本上很多採用這種設計。可以想像如果讓一個 Group 包含更多 Q Head,KV 需要佔用的空間就會更少。

3. Multi-Head Latent Attention (MLA)

既然直接砍 Head 數會讓表達語意的能力打折,那有沒有一種可能:我不砍 Head 數量,但我把 KV 壓縮之後再存? DeepSeek 團隊去年提出了另一種做法 Multi-Head Latent Attention (MLA)。

  • 做法:在把多個 Head 的 K 和 V 存進 Cache 前,先透過一個投影矩陣(這個轉換矩陣是需要在訓練階段學習的),把它們壓縮成維度很小的一個 Latent Vector。在顯存裡只需要存這個被壓縮後的向量!推論階段真正在算 Attention 時,再透過這個 latent 向量 與 訓練好的投影矩陣 還原成原本的 KV 來進行運算。
  • 效果:既能維持原本 Multi-Head 的多視角注意力(智商不掉),又能有效壓縮每個 Token 實際要快取的空間。假設 Latent 向量維度是 512 維,快取只需存這顆向量,1 個 Token 只需要:

https://ithelp.ithome.com.tw/upload/images/20260925/201836073sEhe9bfnZ.png

直接從原本 MHA 的 2 MB 暴砍到現在只剩下 60 KB(256k 只要約 15 GB),節省顯存的效果跟 MQA 一樣顯著,卻幾乎不犧牲 Multi-Head 的優勢。


三、降低 KV Cache 空間的常見策略(二):從 Context 長度下手

剛才我們在算時,都是假設「整整 60 層 Layer,每一層都要把過去 256k 的所有字清楚記下來」。但仔細想想,真的每一層都需要看那麼遠嗎?

事實上,可以只關注較近的 Tokens,因為較近的 Token 也會去參考更前面較遠的 Token,所以在模型後面的 Layer 時,就會把更前面 Tokens 的資訊帶過來。

因此有些模型會採用 Sliding Window Attention 的機制。

  • 以 Google Gemma-4-31B 為例:
    如果再回頭看剛剛那份 gemma-4-31B-it 的 config.json,裡面也有相關設定:
    {
      ...
      "sliding_window": 1024,
      "layer_types": [
      	...
      	"sliding_attention",
      	"sliding_attention",
      	"sliding_attention",
      	"sliding_attention",
      	"sliding_attention",
      	"full_attention",
      	...
      ],
      ...
    }
    
    所以 Gemma-4-31B 採用的策略是,在 60 層的 Layers 中,不是每個 Layer 都是 sliding attention 或都是 full_attention,而是絕大部分 Layer 是 sliding_attention,每隔幾層穿插一層全域的 full_attention。
  • 效果:在那些滑動窗口的層裡面,KV Cache 只要保留最近 1,024 個 Token 的資料就好,超過窗口的舊 Token 在這些 Layer 就不參考。 假設 Gemma-4-31B 的 60 層中,是 50 層 sliding(每層只留 1,024 tokens)配上 10 層 full。 當對話超過 1,024 字後,那 50 層的快取就不會再膨脹了,真正隨長度增加的只剩 10 層。在 GQA 下每一層約 16 KB,原本 60 層全開需要 60 * 16 KB ~= 1 MB / Token,而現在我們可以概算 10 層全局快取,每個 Token 的長期顯存只要:10 Layers * 16 KB ~= 160 KB,直接從原本 1 MB 降到約 160 KB,縮小將近 6 倍。

其實還有很多招,把 FP16 換成 FP8 也是一招。這裡的內容比我想的還要多,感覺今天只能先這樣了,明天再看看怎麼接下去。祝各位中秋節快樂~


參考資料


上一篇
Day 11 | 全知讀者視角:Prompt 從 Agent 到 LLM
下一篇
Day 13 | 明天,我要和昨天的 Token 約會(中):PagedAttention
系列文
在贏家書寫歷史之前:我所看見的 AI,與一位工程師共舞著 共 16 篇
圖片
  熱門推薦
圖片
{{ item.channelVendor }} | {{ item.webinarstarted }} |
{{ formatDate(item.duration) }}
直播中

尚未有邦友留言

立即登入留言