上一篇提到,Decode Attention 的主要成本之一,是反覆讀取前面 tokens 的 K 和 V。
為什麼不能每次都重新算?
因為產生第 t 個 token 時,前面 t-1 個 tokens 的 K、V 都沒有改變。把它們存下來,下次只要計算新 token 的 K、V,再和 cache 裡的內容做 attention。
這就是 KV Cache。
它省下重複計算,代價則是:
Context 越長、requests 越多,GPU memory 就吃得越快。
每個 layer、每個 token 都要保存一份 Key 和一份 Value:
KV Cache Bytes
= Batch
× Sequence Length
× Number of Layers
× 2 # K + V
× Number of KV Heads
× Head Dimension
× Bytes per Element
寫成符號就是:
Memory = B × S × L × 2 × Hkv × Dh × E
其中:
B = batch / 同時存在的 sequences
S = 每條 sequence 的 token 數
L = Transformer layers
Hkv = KV heads
Dh = head dimension
E = 每個元素的 bytes
這是理論 tensor 容量,不包含 allocator、block alignment、page fragmentation 或 quantization scales。
前幾篇使用的 Llama-3.2-1B 參數是:
Layers = 16
Query Heads = 32
KV Heads = 8
Head Dimension = 2048 / 32 = 64
GQA Ratio = 32 / 8 = 4
先用 bf16,也就是每個元素 2 bytes。
每個 token、每個 layer 的 KV Cache:
2 × 8 × 64 × 2 bytes
= 2048 bytes
= 2 KiB
模型有 16 layers,所以每個 token 總共需要:
2 KiB × 16 = 32 KiB
這是一個很好記的數字:
Llama-3.2-1B 的 bf16 GQA KV Cache,每個 token 約 32 KiB。
三者的 Query heads 都維持 32,差別在 KV heads:
MHA:32 個 KV Heads
GQA: 8 個 KV Heads
MQA: 1 個 KV Head
假設其他維度不變,bf16 KV Cache 會是:
| Attention | KV Heads | 每 Token/每 Layer | 每 Token/16 Layers |
|---|---|---|---|
| MHA | 32 | 8 KiB | 128 KiB |
| GQA | 8 | 2 KiB | 32 KiB |
| MQA | 1 | 0.25 KiB | 4 KiB |
因此這個 GQA 4× 設計,KV Cache 是 MHA 的四分之一;MQA 則是 MHA 的三十二分之一。
MHA:每個 Query Head 有自己的 K / V
GQA:一組 Query Heads 共用一個 K / V Head
MQA:所有 Query Heads 共用一個 K / V Head
共享越多,cache 越小;但模型品質、training recipe 和 attention kernel 的設計也會受到影響,所以不是 KV heads 越少就一定越好。
接著把每-token memory 乘上 context length,batch 先固定為 1:
| Context | MHA | GQA(實際模型) | MQA |
|---|---|---|---|
| 2K | 256 MiB | 64 MiB | 8 MiB |
| 8K | 1 GiB | 256 MiB | 32 MiB |
| 32K | 4 GiB | 1 GiB | 128 MiB |
| 128K | 16 GiB | 4 GiB | 512 MiB |
Llama-3.2 支援 128K context。即使只是 1B 模型,單一 128K sequence 的 bf16 GQA KV Cache 就約 4 GiB。
如果使用 MHA,在其他維度相同的假設下會膨脹到 16 GiB。
這也說明為什麼 long-context serving 很快就會從「模型放不放得下」變成「KV Cache 放不放得下」。
固定使用實際的 GQA、bf16、32K context:
| 同時存在的 Sequences | KV Cache |
|---|---|
| 1 | 1 GiB |
| 4 | 4 GiB |
| 16 | 16 GiB |
| 64 | 64 GiB |
這裡的 B 不一定是傳統意義下同時做完的 static batch。
對 continuous batching server 來說,只要 request 還在系統裡、KV blocks 尚未釋放,就會佔用 cache memory。
所以 serving engine 必須管理:
Request 何時進來
→ 分配多少 KV Blocks
→ Context 成長時如何追加
→ Request 完成後何時釋放
這正是後面 PagedAttention 要處理的問題。
固定 GQA、batch 1、32K context:
| KV Cache Format | Bytes/Element | 理論容量 |
|---|---|---|
| FP32 | 4 | 2 GiB |
| BF16 / FP16 | 2 | 1 GiB |
| FP8 | 1 | 512 MiB |
| 4-bit(理想值) | 0.5 | 256 MiB |
位元數減半,理論 tensor 容量也減半。
但實際 quantized KV Cache 還可能需要 scale、zero point、group metadata、alignment 和 padding,因此不一定精確等於表中的理想值。
而且壓縮 KV Cache 不只是在省容量。Decode 每一步都要讀取歷史 K、V;cache 變小,也可能降低 HBM traffic。
代價則是多了 quantize / dequantize,以及可能的 accuracy loss。