KV Cache:為什麼推理需要它

從自回歸生成的重複計算到記憶體瓶頸,拆解 KV Cache 的原理、大小計算、優化技術與最佳實踐。

KV Cache:為什麼推理需要它

上一篇我們談了 LLM 部署的三大指標:延遲、吞吐、成本。

其中,延遲的關鍵瓶頸之一,就是自回歸生成——模型必須一個一個 token 地產出,每一步都依賴前一步的結果。

如果沒有優化,這個過程會極度浪費算力:每生成一個新 token,都要重新計算所有先前 token 的注意力。

KV Cache 就是為了解決這個問題而生的技術。它是 LLM 推理最核心的優化機制,也是記憶體瓶頸的根源。

這一篇,我們要深入 KV Cache 的原理、計算、瓶頸,以及各種優化技術。

Self-Attention → KV Cache → Memory Bottleneck

一、先理解問題:自回歸生成的重複計算

沒有 KV Cache 會怎樣?

假設模型要生成「今天天氣真好」這五個字。

  • 生成「今」的時候:模型看到輸入,計算注意力,產出「今」。
  • 生成「天」的時候:模型看到「今」,需要重新計算「今」的 Key 和 Value,再產出「天」。
  • 生成「天」的時候:模型看到「今天」,需要重新計算「今」和「天」的 Key 和 Value,再產出「天」。
  • 生成「氣」的時候:模型看到「今天天」,需要重新計算「今」、「天」、「天」的 Key 和 Value……

你有沒有發現問題?

每生成一個新 token,模型都要重新計算所有先前 token 的 Key 和 Value。這是大量的重複計算。

用具體數字來看:

生成步驟需要計算的 token 數累計計算量
第 1 個11
第 2 個23
第 3 個36
第 4 個410
第 5 個515
第 N 個NN(N+1)/2

如果生成 1000 個 token,總計算量是 500,500,而不是 1000。
這是 O(N²) 的複雜度,極度浪費。

為什麼會重複計算?

回顧注意力機制的公式:

Attention(Q,K,V)=softmax(QKTdk)V\text{Attention}(Q, K, V) = \text{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)V
  • Q(Query):當前 token 的查詢
  • K(Key):所有 token 的鍵
  • V(Value):所有 token 的值

當生成新 token 時:

  1. 新的 Q 需要與所有 K 計算注意力。
  2. 但先前 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 個11
第 2 個21
第 3 個31
第 4 個41
第 5 個51
第 N 個N1
總計N(N+1)/2N

複雜度從 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 是推理的關鍵優化,但它有一個致命的代價:記憶體

計算公式

KV Cache 大小=2×L×H×D×S×B×P\text{KV Cache 大小} = 2 \times L \times H \times D \times S \times B \times P
符號意義範例(70B 模型)
2Key 和 Value 各一份2
L層數80
H注意力頭數64
D每個頭的維度128
S序列長度4096
B批次大小32
P精度(bytes)2(FP16)

具體計算

以 LLaMA-2 70B 為例:

KV Cache 大小=2×80×64×128×4096×32×2 bytes=343,597,383,680 bytes320 GiB\begin{aligned} \text{KV Cache 大小} &= 2 \times 80 \times 64 \times 128 \times 4096 \times 32 \times 2 \text{ bytes} \\ &= 343,597,383,680 \text{ bytes} \\ &\approx 320 \text{ GiB} \end{aligned}

320 GB 的 KV Cache!
這遠超過單張 A100 的 80 GB 記憶體。

不同模型的 KV Cache 大小

模型層數頭數頭維度序列長度 4096,批次 1 的 KV Cache
GPT-212126412 MB
LLaMA-7B323212864 MB
LLaMA-13B4040128100 MB
LLaMA-70B8064128320 MB
GPT-3 175B9696128576 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 的大小,直接限制了批次大小,進而限制了吞吐。

Model Weights + KV Cache < GPU Memory

四、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 大小品質
MHA6464320 MB最高
GQA(G=8)64840 MB接近 MHA
MQA6415 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 大小縮小倍數
FP16320 MB1x
INT8160 MB2x
INT480 MB4x

挑戰:精度損失、需要校準縮放因子。

技術五: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 縮小品質影響實作難度
MQAH 倍
GQAH/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 會是:

7.5 MB×(409615)2048 MB=2 GB7.5 \text{ MB} \times \left(\frac{4096}{15}\right) \approx 2048 \text{ MB} = 2 \text{ GB}

這只是一個請求。如果是批次 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 自動處理這個優化。

七、常見的陷阱

  1. 忘記啟用 KV Cacheoutputs = model(input_ids) 預設可能沒啟用,要明確加 use_cache=True
  2. 忽略 KV Cache 的大小:只計算模型權重,忽略 KV Cache,導致記憶體不足。
  3. 序列太長:用 32K 的上下文,KV Cache 爆掉。只保留必要的上下文。
  4. 批次太大:批次 128,KV Cache 爆掉。根據記憶體計算最大批次。
  5. 沒有用 GQA:用 MHA 模型,KV Cache 大。改用 GQA 模型。
  6. 沒有用 PagedAttention:用預設推理,改用 vLLM。
  7. 沒有監控:不知道 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),以及如何在實際部署中選擇合適的量化策略。