KV Cache:為什麼推理需要它
從自回歸生成的重複計算到記憶體瓶頸,拆解 KV Cache 的原理、大小計算、優化技術與最佳實踐。
KV Cache:為什麼推理需要它
上一篇我們談了 LLM 部署的三大指標:延遲、吞吐、成本。
其中,延遲的關鍵瓶頸之一,就是自回歸生成——模型必須一個一個 token 地產出,每一步都依賴前一步的結果。
如果沒有優化,這個過程會極度浪費算力:每生成一個新 token,都要重新計算所有先前 token 的注意力。
KV Cache 就是為了解決這個問題而生的技術。它是 LLM 推理最核心的優化機制,也是記憶體瓶頸的根源。
這一篇,我們要深入 KV Cache 的原理、計算、瓶頸,以及各種優化技術。
一、先理解問題:自回歸生成的重複計算
沒有 KV Cache 會怎樣?
假設模型要生成「今天天氣真好」這五個字。
- 生成「今」的時候:模型看到輸入,計算注意力,產出「今」。
- 生成「天」的時候:模型看到「今」,需要重新計算「今」的 Key 和 Value,再產出「天」。
- 生成「天」的時候:模型看到「今天」,需要重新計算「今」和「天」的 Key 和 Value,再產出「天」。
- 生成「氣」的時候:模型看到「今天天」,需要重新計算「今」、「天」、「天」的 Key 和 Value……
你有沒有發現問題?
每生成一個新 token,模型都要重新計算所有先前 token 的 Key 和 Value。這是大量的重複計算。
用具體數字來看:
| 生成步驟 | 需要計算的 token 數 | 累計計算量 |
|---|---|---|
| 第 1 個 | 1 | 1 |
| 第 2 個 | 2 | 3 |
| 第 3 個 | 3 | 6 |
| 第 4 個 | 4 | 10 |
| 第 5 個 | 5 | 15 |
| … | … | … |
| 第 N 個 | N | N(N+1)/2 |
如果生成 1000 個 token,總計算量是 500,500,而不是 1000。
這是 O(N²) 的複雜度,極度浪費。
為什麼會重複計算?
回顧注意力機制的公式:
- Q(Query):當前 token 的查詢
- K(Key):所有 token 的鍵
- V(Value):所有 token 的值
當生成新 token 時:
- 新的 Q 需要與所有 K 計算注意力。
- 但先前 token 的 K 和 V 沒有改變。
既然沒有改變,為什麼要重新計算?
這就是 KV Cache 的核心洞察。
二、KV Cache 的原理
核心思想
把先前 token 的 Key 和 Value 快取起來,生成新 token 時直接使用,避免重複計算。
用一個比喻:
- 沒有 KV Cache:每次做菜都要重新切所有食材。
- 有 KV Cache:第一次切好的食材放冰箱,下次直接拿出來用。
運作流程
預填充階段(Prefill):
輸入:「今天天氣」
1. 計算所有 token 的 Q、K、V
2. 將 K、V 存入 KV Cache
3. 產出第一個輸出 token
解碼階段(Decode):
生成「真」:
1. 只計算新 token 的 Q
2. 從 KV Cache 取出所有先前的 K、V
3. 計算注意力
4. 產出「真」
5. 將「真」的 K、V 加入 KV Cache
生成「好」:
1. 只計算「好」的 Q
2. 從 KV Cache 取出所有先前的 K、V(今天天氣真)
3. 計算注意力
4. 產出「好」
5. 將「好」的 K、V 加入 KV Cache
計算量的對比
| 生成步驟 | 沒有 KV Cache | 有 KV Cache |
|---|---|---|
| 第 1 個 | 1 | 1 |
| 第 2 個 | 2 | 1 |
| 第 3 個 | 3 | 1 |
| 第 4 個 | 4 | 1 |
| 第 5 個 | 5 | 1 |
| … | … | … |
| 第 N 個 | N | 1 |
| 總計 | N(N+1)/2 | N |
複雜度從 O(N²) 降到 O(N)。
具體的程式碼對比
沒有 KV Cache:
def generate_without_cache(model, input_ids, max_new_tokens):
"""沒有 KV Cache 的生成"""
generated = input_ids
for _ in range(max_new_tokens):
# 每次都重新計算所有 token 的 K、V
outputs = model(generated)
next_token = sample(outputs.logits[:, -1, :])
generated = torch.cat([generated, next_token], dim=-1)
return generated
有 KV Cache:
def generate_with_cache(model, input_ids, max_new_tokens):
"""有 KV Cache 的生成"""
generated = input_ids
past_key_values = None
for _ in range(max_new_tokens):
# 第一次傳入完整輸入,之後只傳入新 token
if past_key_values is None:
outputs = model(generated, use_cache=True)
else:
outputs = model(
generated[:, -1:], # 只傳入最後一個 token
past_key_values=past_key_values,
use_cache=True,
)
# 更新 KV Cache
past_key_values = outputs.past_key_values
next_token = sample(outputs.logits[:, -1, :])
generated = torch.cat([generated, next_token], dim=-1)
return generated
第二個版本只傳入新 token,KV Cache 自動處理先前的 K、V。
這就是為什麼現代 LLM 推理框架都預設啟用 KV Cache。
三、KV Cache 的大小計算
KV Cache 是推理的關鍵優化,但它有一個致命的代價:記憶體。
計算公式
| 符號 | 意義 | 範例(70B 模型) |
|---|---|---|
| 2 | Key 和 Value 各一份 | 2 |
| L | 層數 | 80 |
| H | 注意力頭數 | 64 |
| D | 每個頭的維度 | 128 |
| S | 序列長度 | 4096 |
| B | 批次大小 | 32 |
| P | 精度(bytes) | 2(FP16) |
具體計算
以 LLaMA-2 70B 為例:
320 GB 的 KV Cache!
這遠超過單張 A100 的 80 GB 記憶體。
不同模型的 KV Cache 大小
| 模型 | 層數 | 頭數 | 頭維度 | 序列長度 4096,批次 1 的 KV Cache |
|---|---|---|---|---|
| GPT-2 | 12 | 12 | 64 | 12 MB |
| LLaMA-7B | 32 | 32 | 128 | 64 MB |
| LLaMA-13B | 40 | 40 | 128 | 100 MB |
| LLaMA-70B | 80 | 64 | 128 | 320 MB |
| GPT-3 175B | 96 | 96 | 128 | 576 MB |
注意:這是批次 1、序列長度 4096 的大小。
如果批次是 32,大小要乘以 32。
為什麼 KV Cache 是記憶體瓶頸?
一張 A100 有 80 GB 記憶體。
如果模型權重佔 140 GB(70B,FP16),需要兩張 A100。
剩下的記憶體用來放 KV Cache。
兩張 A100:160 GB
模型權重:140 GB
可用於 KV Cache:20 GB
每個請求的 KV Cache:320 MB(序列長度 4096)
最大批次:20 GB / 320 MB ≈ 64 個請求
如果序列長度增加到 8192,每個請求的 KV Cache 變成 640 MB,最大批次降到 32。
KV Cache 的大小,直接限制了批次大小,進而限制了吞吐。
四、KV Cache 的優化技術
既然 KV Cache 是瓶頸,就有很多優化技術。
技術一:MQA(Multi-Query Attention)
核心思想:讓所有注意力頭共享同一組 Key 和 Value。
標準的多頭注意力(MHA)中,每個頭有自己的 Q、K、V:
KV Cache 大小:2 × L × H × D × S × B × P
MQA 中,所有頭共享同一組 K、V:
KV Cache 大小:2 × L × 1 × D × S × B × P
KV Cache 縮小 H 倍(頭數)。
以 LLaMA-70B 為例,H = 64:
MHA:320 MB
MQA:5 MB
縮小 64 倍!
但 MQA 有個缺點:品質下降。
所有頭共享 K、V,失去了多頭的多樣性。
技術二:GQA(Grouped-Query Attention)
核心思想:折衷方案。把頭分成幾組,每組共享一組 K、V。
- 標準 MHA:每個頭有自己的 K、V(H 組)
- GQA:每 G 個頭共享一組 K、V(H/G 組)
- MQA:所有頭共享一組 K、V(1 組)
以 LLaMA-70B 為例:
| 方法 | 頭數 | KV 頭數 | KV Cache 大小 | 品質 |
|---|---|---|---|---|
| MHA | 64 | 64 | 320 MB | 最高 |
| GQA(G=8) | 64 | 8 | 40 MB | 接近 MHA |
| MQA | 64 | 1 | 5 MB | 略降 |
GQA 是目前的業界標準。
LLaMA-2 70B、LLaMA-3、Mistral 等模型都採用 GQA。
class GroupedQueryAttention(nn.Module):
"""GQA 實作"""
def __init__(self, d_model, n_heads, n_kv_heads):
super().__init__()
self.n_heads = n_heads
self.n_kv_heads = n_kv_heads
self.head_dim = d_model // n_heads
self.q_proj = nn.Linear(d_model, n_heads * self.head_dim)
self.k_proj = nn.Linear(d_model, n_kv_heads * self.head_dim)
self.v_proj = nn.Linear(d_model, n_kv_heads * self.head_dim)
self.o_proj = nn.Linear(d_model, d_model)
def forward(self, x, past_kv=None):
batch, seq_len, _ = x.shape
q = self.q_proj(x).view(batch, seq_len, self.n_heads, self.head_dim)
k = self.k_proj(x).view(batch, seq_len, self.n_kv_heads, self.head_dim)
v = self.v_proj(x).view(batch, seq_len, self.n_kv_heads, self.head_dim)
# 重複 K、V 以匹配 Q 的頭數
k = k.repeat_interleave(self.n_heads // self.n_kv_heads, dim=2)
v = v.repeat_interleave(self.n_heads // self.n_kv_heads, dim=2)
# 注意力計算...
技術三:PagedAttention
核心思想:把 KV Cache 分成固定大小的「頁」,像作業系統的虛擬記憶體一樣管理。
傳統 KV Cache 的問題:
- 碎片化:每個請求的 KV Cache 大小不同,造成記憶體碎片。
- 浪費:預留最大長度,但實際可能用不到。
- 無法共享:相同的前綴無法共享 KV Cache。
PagedAttention 的做法:
傳統:
[請求 A 的 KV Cache(連續)][請求 B 的 KV Cache(連續)]
問題:A 和 B 之間可能有碎片,浪費記憶體
Paged:
[頁 1][頁 2][頁 3][頁 4][頁 5][頁 6]
A.1 B.1 A.2 B.2 A.3 B.3
優點:沒有碎片,記憶體利用率接近 100%
PagedAttention 的優點:
| 優點 | 說明 |
|---|---|
| 減少碎片 | 記憶體利用率從 20-40% 提升到 90%+ |
| 共享前綴 | 相同的系統提示可以共享 KV Cache |
| 動態分配 | 根據實際需要分配頁 |
PagedAttention 是 vLLM 的核心技術。
技術四:KV Cache 量化
核心思想:把 KV Cache 從 FP16 量化到 INT8 或 INT4。
FP16:2 bytes per element
INT8:1 byte per element(縮小 2 倍)
INT4:0.5 bytes per element(縮小 4 倍)
以 LLaMA-70B 為例:
| 精度 | KV Cache 大小 | 縮小倍數 |
|---|---|---|
| FP16 | 320 MB | 1x |
| INT8 | 160 MB | 2x |
| INT4 | 80 MB | 4x |
挑戰:精度損失、需要校準縮放因子。
技術五:KV Cache 壓縮
核心思想:不是所有 token 的 KV 都同等重要。壓縮或丟棄不重要的。
常見方法:
| 方法 | 說明 |
|---|---|
| H2O | 保留注意力分數高的 token |
| StreamingLLM | 保留開頭的幾個 token 和最近的 token |
| SnapKV | 根據注意力模式選擇重要的 token |
這些方法可以在幾乎不損失品質的情況下,大幅減少 KV Cache 大小。
技術六:滑動窗口注意力
核心思想:只關注最近的 N 個 token,而不是所有 token。
標準注意力:token i 關注所有 token 1 到 i
滑動窗口:token i 只關注 token i-W 到 i(W 是窗口大小)
這樣 KV Cache 只需要保留最近 W 個 token。
Mistral 7B 就採用了滑動窗口注意力,窗口大小 4096。
技術對比
| 技術 | KV Cache 縮小 | 品質影響 | 實作難度 |
|---|---|---|---|
| MQA | H 倍 | 中 | 低 |
| GQA | H/G 倍 | 低 | 低 |
| PagedAttention | 不縮小,但提升利用率 | 無 | 中 |
| KV 量化 | 2-4 倍 | 低 | 中 |
| KV 壓縮 | 2-10 倍 | 中 | 高 |
| 滑動窗口 | 固定大小 | 中 | 低 |
在實際部署中,這些技術通常組合使用:
GQA(縮小 KV Cache)+ PagedAttention(提升利用率)+ INT8 量化(再縮小 2 倍)
五、實作:觀察 KV Cache 的行為
讓我們用一個具體的例子,觀察 KV Cache 的行為。
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer
def analyze_kv_cache(model_name: str = "meta-llama/Llama-2-7b-hf"):
"""分析 KV Cache 的行為"""
tokenizer = AutoTokenizer.from_pretrained(model_name)
model = AutoModelForCausalLM.from_pretrained(
model_name,
torch_dtype=torch.float16,
device_map="auto",
)
prompt = "今天天氣真好,適合"
inputs = tokenizer(prompt, return_tensors="pt").to(model.device)
with torch.no_grad():
outputs = model.generate(
**inputs,
max_new_tokens=10,
use_cache=True,
return_dict_in_generate=True,
output_scores=True,
)
# 觀察 KV Cache 的形狀
print("KV Cache 結構:")
print(f"層數:{len(outputs.past_key_values)}")
print(f"每層的 K、V 形狀:{outputs.past_key_values[0][0].shape}")
# 計算 KV Cache 大小
total_bytes = 0
for layer_kv in outputs.past_key_values:
for tensor in layer_kv:
total_bytes += tensor.numel() * tensor.element_size()
print(f"\nKV Cache 總大小:{total_bytes / 1024 / 1024:.2f} MB")
if __name__ == "__main__":
analyze_kv_cache()
執行結果
KV Cache 結構:
層數:32
每層的 K、V 形狀:torch.Size([1, 32, 15, 128])
KV Cache 總大小:7.50 MB
即使只有 15 個 token,KV Cache 已經佔了 7.5 MB。
如果序列長度是 4096,KV Cache 會是:
這只是一個請求。如果是批次 32,就是 64 GB。
六、KV Cache 的最佳實踐
1. 總是啟用 KV Cache
# 不好的例子:沒有 KV Cache
outputs = model(input_ids)
# 好的例子:啟用 KV Cache
outputs = model(input_ids, use_cache=True)
2. 選擇支援 GQA 的模型
# 不好的例子:用 MHA 模型(KV Cache 大)
# 好的例子:用 GQA 模型(KV Cache 小)
# LLaMA-3、Mistral、Qwen 等都支援 GQA
3. 使用 PagedAttention
# 不好的例子:用 HuggingFace 的預設推理
# 好的例子:用 vLLM(內建 PagedAttention)
from vllm import LLM
llm = LLM(model="meta-llama/Llama-2-7b-hf")
outputs = llm.generate(prompts)
4. 控制序列長度
# 不好的例子:不限制 max_tokens
# 好的例子:根據需求設定合理的上限
outputs = model.generate(
input_ids,
max_new_tokens=200, # 不要設太大
)
5. 監控 KV Cache 使用率
def monitor_kv_cache(model):
"""監控 KV Cache 使用率"""
if hasattr(model, "cache"):
cache_size = sum(
t.numel() * t.element_size()
for layer in model.cache
for t in layer
)
print(f"KV Cache 大小:{cache_size / 1024**3:.2f} GB")
6. 使用量化
model = AutoModelForCausalLM.from_pretrained(
model_name,
load_in_8bit=True, # 8-bit 量化
)
7. 批次大小要合理
# 根據 KV Cache 大小計算最大批次
max_batch = available_memory // kv_cache_per_request
8. 共享前綴
如果多個請求有相同的前綴(如系統提示),可以共享 KV Cache。vLLM 自動處理這個優化。
七、常見的陷阱
- 忘記啟用 KV Cache:
outputs = model(input_ids)預設可能沒啟用,要明確加use_cache=True。 - 忽略 KV Cache 的大小:只計算模型權重,忽略 KV Cache,導致記憶體不足。
- 序列太長:用 32K 的上下文,KV Cache 爆掉。只保留必要的上下文。
- 批次太大:批次 128,KV Cache 爆掉。根據記憶體計算最大批次。
- 沒有用 GQA:用 MHA 模型,KV Cache 大。改用 GQA 模型。
- 沒有用 PagedAttention:用預設推理,改用 vLLM。
- 沒有監控:不知道 KV Cache 用了多少,持續監控。
八、總結:KV Cache 是必要的,也是昂貴的
讓我們回顧這一篇的核心:
- 為什麼需要 KV Cache:自回歸生成會重複計算先前的 K、V,複雜度 O(N²)。
- KV Cache 的原理:把先前的 K、V 快取起來,複雜度降到 O(N)。
- KV Cache 的大小:與層數、頭數、頭維度、序列長度、批次大小成正比。
- KV Cache 是記憶體瓶頸:70B 模型、序列 4096、批次 32 需要 320 GB。
- 優化技術:MQA、GQA、PagedAttention、量化、壓縮、滑動窗口。
- 最佳實踐:總是啟用、選擇 GQA 模型、使用 PagedAttention、控制序列長度、監控、量化、合理批次、共享前綴。
- 常見陷阱:忘記啟用、忽略大小、序列太長、批次太大、沒有用 GQA、沒有用 PagedAttention、沒有監控。
KV Cache 是 LLM 推理的必要機制。
沒有它,推理會慢得無法使用。
有了它,記憶體成為新的瓶頸。
理解了 KV Cache,你就理解了為什麼 LLM 部署需要那麼多優化技術。
接下來,我們要深入第一個優化技術:量化。
下一篇預告
《量化:FP16、INT8、INT4 的取捨》
我們會解釋量化的原理、不同精度的差異、量化對延遲與成本的影響、常見的量化方法(GPTQ、AWQ、GGUF),以及如何在實際部署中選擇合適的量化策略。