iT邦幫忙

2026 iThome 鐵人賽

DAY 11
0
Software Development

LLM infra 學習日記系列 第 11

KV Cache 到底有多大?手算 MHA、GQA 與 MQA

  • 分享至 

  • xImage
  •  

上一篇提到,Decode Attention 的主要成本之一,是反覆讀取前面 tokens 的 K 和 V。

為什麼不能每次都重新算?

因為產生第 t 個 token 時,前面 t-1 個 tokens 的 K、V 都沒有改變。把它們存下來,下次只要計算新 token 的 K、V,再和 cache 裡的內容做 attention。

這就是 KV Cache。

它省下重複計算,代價則是:

Context 越長、requests 越多,GPU memory 就吃得越快。


KV Cache 的公式

每個 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 來算

前幾篇使用的 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。


MHA、GQA、MQA 差多少?

三者的 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 越少就一定越好。


Context Length 放大後

接著把每-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 放不放得下」。


Batch 與 Concurrency 是直接相乘

固定使用實際的 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 要處理的問題。


Dtype 能省多少?

固定 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。


上一篇
FlashAttention 解決了什麼?從 Tiling 到 Online Softmax
下一篇
PagedAttention:為什麼 KV Cache 也需要分頁?
系列文
LLM infra 學習日記19
圖片
  熱門推薦
圖片
{{ item.channelVendor }} | {{ item.webinarstarted }} |
{{ formatDate(item.duration) }}
直播中

尚未有邦友留言

立即登入留言