← Back to list

Transformer 雷之呼吸第一型:KV Cache

從 KV Cache 到 MLA,探討讓大語言模型解碼加速數十倍的方法

Rice Yang · 2025-08-17 09:42 · 1 claps · 15.4 min read
#llm #kv-cache #mqa #aml #transformer-model
Open on Medium ↗
Wiki topics: LLM · Large Language Models OPS · LLMOps & Inference

Transformer 雷之呼吸第一型:KV Cache

從 KV Cache 到 MLA,探討讓大語言模型解碼加速數十倍的方法

來源:由 GPT-5 生成。我要求他生成一張雷之呼吸浮世繪風格的 AI 封面圖

來源:由 GPT-5 生成。我要求他生成一張雷之呼吸浮世繪風格的 AI 封面圖

如今,Transformer 已成為大型語言模型(LLM)和基礎 AI 模型的基本範式。現在,Transformer 不僅在跟你聊天,還在解決你的實際問題。它需要理解你問題的更多上下文,例如閱讀 900 頁的大而美法案以判斷對你業務的可能影響,或理解 10 萬行程式碼來修復軟體中的隱藏錯誤。

然而,這些模型的輸入長度有限,長上下文輸入與推理成為真正影響 AI 產業的實際技術挑戰。有許多技術方向可以解決長上下文問題。解決長上下文問題有兩個好處,一個是可以讓模型讀得更長、寫得更長,另一個是可以讓模型變得更快。

本文試圖深入其中的一個方向:KV Cache。KV Cache 是 LLM 部署中最好實現的優化之一,可以讓一個 LLM 提升 10 倍以上的推理速度與上下文長度。為了突顯它的加速效果,我戲稱為 Transformer 的雷之呼吸第一型。

QKV in Self‑Attention

我們在《*The Great Transformer”》中介紹過 Self‑Attention(SA)*,它是 Transformer 的關鍵機制。在 SA 中,每個輸入 token(實際上代表 LLM 中的一個詞、像素或概念)會被投影成三個向量:Query、Key 和 Value(QKV),這些向量是 SA 演算法的輸入。

來源: KV Cache Explained Intuitively

來源: KV Cache Explained Intuitively

當我們增加更多 tokens,也就是向 AI 模型添加更多詞、像素或概念時,該 token 會從向量擴展為矩陣。新的 token 會被添加到矩陣的最後一行。例如,當我們將 tokens 從 1 增加到 3 時,看起來會像這樣:

來源: KV Cache Explained Intuitively

來源: KV Cache Explained Intuitively

儘管在 Transformer 推理中,softmax 計算通常是主要瓶頸,但從 token 到 Q/K/V 的投影計算也很重,尤其是在解碼器 decoder 中。

Decoder Inference 解碼器推理

假設你的 LLM 有 L 層,輸出長度為 N,則 decoder 簡化的推理流程如下:

  1. 將輸入 X(長度=1)送入 L 層的 LLM。每一層產生 1 個 QKV 向量。decoder 生成一個 token y₁。
  2. y₁ 串接到輸入 X,然後將 X(長度=2)送入 L 層的 LLM。每一層產生 2 個 QKV 向量。decoder 生成一個 token y₂。
  3. y₂ 串接到輸入 X,然後將 X(長度=3)送入 L 層的 LLM。每一層產生 3 個 QKV 向量。decoder 生成一個 token y₃。
  4. 重複以上步驟
  5. yₙ₋₁ 串接到輸入 X,然後將 X(長度=n)送入 L 層的 LLM。每一層產生 n 個 QKV 向量。decoder 生成一個 token yₙ。

Decoder 需要多次迭代才能解碼每一個單詞。來源:The Illustrated Transformer

Decoder 需要多次迭代才能解碼每一個單詞。來源:The Illustrated Transformer

這個過程的問題在於我們需要計算 L × (1 + 2 + … + N) 次 QKV,其時間複雜度為 O(N²)。好消息是,我們可以透過快取來降低複雜度。

KV Cache

如果我們細看 SA 公式,每個 token 位置,也就是輸入 Q、K、V 和輸出 Z 矩陣中的每一行,都與所有其他位置相關聯 — — 這就是 Self‑Attention 的設計。

自注意力公式的矩陣計算。來源: The Illustrated Transformer

自注意力公式的矩陣計算。來源: The Illustrated Transformer

但在 decoder 中,每一行 Q、K 和 V 只與其歷史相關,與未來無關。在更通用的機器學習用語中,我們稱為 *autoregression(自回歸)。因此,LLM 的 decoder 在 softmax 中設計了 [causal mask](https://medium.com/@jinoo/a-simple-example-of-attention-masking-in-transformer-decoder-a6c66757bc7d)*(因果遮罩)以遮蔽未來

在這種情況下,第 i 行的 QKV 只與編號小於 i 的行相關。為了達到這一點,我們只需將第 i 行的 Key (K)Value (V) 串接到當前第 i‑1 行的 KV 底部。

左:因果遮罩causal mask。右:經過因果遮罩計算的注意力矩陣 Attention matrix

左:因果遮罩causal mask。右:經過因果遮罩計算的注意力矩陣 Attention matrix

為了實驗這樣簡單的串接 (concatinate) 並不會破壞 decoder 中 SA 層的計算,我寫了以下的 Python 實驗程式碼:

import torch
import math
import matplotlib.pyplot as plt

torch.set_printoptions(precision=6, sci_mode=False)
# ----- Hyperparameters -----
seed = 0
torch.manual_seed(seed)
d_model = 1   # single value (1x1) as requested
d_k = 1       # single-head attention with 1-dim projections
L = 3         # number of layers
T = 5         # context length to grow to
d_ff = 4      # hidden size for the feed-forward layer
dtype = torch.float64
device = torch.device("cpu")
# ----- Random weights and initial token -----
WQ = torch.randn(d_model, d_k, dtype=dtype, device=device)
WK = torch.randn(d_model, d_k, dtype=dtype, device=device)
WV = torch.randn(d_model, d_k, dtype=dtype, device=device)
# Feed-forward layer weights (shared across layers)
W1 = torch.randn(d_k, d_ff, dtype=dtype, device=device)
W2 = torch.randn(d_ff, d_model, dtype=dtype, device=device)
X = torch.randn(1, d_model, dtype=dtype, device=device)  # initial (1x1)

def causal_self_attention(X):
    """
    X: (t, d_model)
    Returns: (out, Q, K, V, attn)
    """
    Q = X @ WQ  # (t, d_k)
    K = X @ WK  # (t, d_k)
    V = X @ WV  # (t, d_k)
    t = X.size(0)
    scores = (Q @ K.T) / math.sqrt(d_k)  # (t, t)
    mask = torch.triu(torch.ones(t, t, dtype=torch.bool, device=X.device), diagonal=1)
    scores = scores.masked_fill(mask, -1e9)
    attn = torch.softmax(scores, dim=-1)  # (t, t)
    out = attn @ V                        # (t, d_k)
    return out, Q, K, V, attn, mask

def feed_forward(X):
    """Position-wise feed-forward: ReLU(X @ W1) @ W2
    X: (t, d_k) -> (t, d_model)
    """
    return torch.relu(X @ W1) @ W2
print("=== Weights ===")
print("WQ:", WQ.view(-1))
print("WK:", WK.view(-1))
print("WV:", WV.view(-1))
print("W1:", W1.view(-1))
print("W2:", W2.view(-1))
print("\nInitial X (t=1):", X.view(-1))
for t in range(1, T + 1):
    X_l = X
    Ks, Vs, Qs = [], [], []
    # Apply the SAME SA L times (L identical layers), followed by FF each time
    for l in range(L):
        Z, Q, K, V, attn, mask = causal_self_attention(X_l)
        # Feed-forward layer after self-attention
        X_l = feed_forward(Z)  # (t, d_model) -> next layer input
    # Report (use the last layer to keep output compact)
    print(f"\n=== Step t={t} (sequence length {X.size(0)}) ===")
    print(f"K (layer {L}) shape={tuple(K.shape)}:\n{K.view(-1)}")
    print(f"V (layer {L}) shape={tuple(V.shape)}:\n{V.view(-1)}")
    print(f"Q (layer {L}) shape={tuple(Q.shape)}:\n{Q.view(-1)}")
    print(f"Z (layer {L}) shape={tuple(Z.shape)}:\n{Z.view(-1)}")
    # ----- Visualization -----
    fig, axs = plt.subplots(1, 3, figsize=(15, 5))
    # Q, K, V bar plots
    axs[0].bar(range(Q.shape[0]), Q.view(-1).cpu().numpy(), label="Q")
    axs[0].set_title('Q')
    axs[1].bar(range(K.shape[0]), K.view(-1).cpu().numpy(), label="K", color='orange')
    axs[1].set_title('K')
    axs[2].bar(range(V.shape[0]), V.view(-1).cpu().numpy(), label="V", color='green')
    axs[2].set_title('V')
    # Mask as image
    # Create a new figure for mask and attn if t > 3 to prevent small images
    fig2, axs2 = plt.subplots(1, 2, figsize=(8, 4))
    axs2[0].imshow(mask.cpu().numpy(), cmap='gray', vmin=0, vmax=1)
    axs2[0].set_title('Mask')
    axs2[1].imshow(attn.detach().cpu().numpy(), cmap='viridis')
    axs2[1].set_title('Attention')
    plt.tight_layout()
    plt.show()
    # Grow sequence by appending the LAST token of the final layer's output
    if t < T:
        new_token = X_l[-1:].detach()   # (1, d_model)
        X = torch.cat([X, new_token], dim=0)
print("\nDone. Observation: For every t>1 and for every layer, the first t-1 rows of K and V "
      "are exactly equal to those computed at t-1. Hence, KV is cacheable.")

這段程式碼的輸出如下:

第 1 層 QKV

第 1 層 QKV

第 2 層 QKV

第 2 層 QKV

第 3 層 QKV

第 3 層 QKV

第 4 層 QKV

第 4 層 QKV

第 5 層 QKV

第 5 層 QKV

因此實際上,我們不需要重新計算舊的 KV,只需要快取歷史的 KV,並將新的向量 k 和 v 串接上去就行了。實驗結果也證明這樣的操作不會破壞本來的 KV 計算。這將把時間複雜度從 O(N²) 降至 O(N)。我們把這個技巧叫作 KV cache

這裡可能有一個疑問:既然 V 也用到了 causal mask,也跟 KV 一樣可以被串接,為什麼我們只快取 KV,而不快取 Q 呢?因為第 i 個輸出只依賴於 QV 矩陣和第 iQuery 向量,所以快取過去的 V 矩陣沒有意義,我們後面也不會再用到它。

來源: KV Cache Secrets: Boost LLM Inference Efficiency

來源: KV Cache Secrets: Boost LLM Inference Efficiency

但我們還能再更快嗎?我們已經對 Self‑Attention 層進行了優化,現在可以試試對多頭注意力 (Multi‑Head Attention, MHA)搞點東西。

MQA 和 GQA

為了加速,我們需要壓縮資訊。在 MHA 中,我們有多個自注意力分支(稱為「頭 (head) 」),最後將它們合併。每個 head 都有自己的 QKV 矩陣。一個直接的方案是:共享各個 headKV

來源: GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints

來源: GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints

插圖中每一列代表一個「head」。最左邊是最原始的 Multi‑Head Attention ,每個頭有各自的 QKV

最右邊的 Multi‑Query Attention (MQA)是最激進的優化方法,所有頭之間共享 KVMQA 能節省最多的記憶體和投影計算,但由於跨頭壓縮了 K/V 資訊,也必然會降低模型的精度。

中間的 Grouped‑Query Attention (GQA)Google Research 在 2023 年提出的一篇簡潔優雅的論文。有些頭共享同一組 VK,但保留各自的 Query。這在 MHAMQA 之間取得了良好的平衡:既保持接近 MHA 的精度,又使推論速度非常接近 MQA。下圖展示了 MHAMQAGQA 在速度與精度之間的權衡 (speed-accuracy trade-off)。

來源: GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints

來源: GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints

順便一提,他們只共享 KV 而不共享 Q 的原因是:如果三者都共享,那每個頭的輸出都會一模一樣,失去 multi‑head 的意義。

Google 讓 K/V 計算更簡單,而 DeepSeek 往更深的方向改進。

The Multi-Latent Attention, MLA

在帶 KV cacheMHA decoder 中,將每個輸入 token 投影為 QKV 的時間複雜度為 O(H × L × dₖ× d),其中 H 是頭數,L 是當前輸入長度,dₖQ/K/V 每列的維度,d 是輸入 token 的維度。

使用 MQA 或 GQA 可以透過減少冗餘計算或數據壓縮提升速度,但時間複雜度本身並不會變,因為投影 V 還是不可避免計算的。

*DeepSeek V2 在論文中提出了 Multi‑Latent Attention (MLA),將 K/V 投影到維度為 r 的隱空間 (latent space),其中 rdₖ。當我們從 d 維投影到 dₖ* 維時,可以先投影到更低維度,再重新投影回高維度。

來源: DeepSeek V2 paper

來源: DeepSeek V2 paper

舉例來說,若我們將 d 維輸入 token 投影為 dₖ 維的向量,計算時間是 d × dₖ。但如果先投影到另一個維度 r,再重新投影回 dₖ,計算時間變為 (d × r) + (dₖ × r)

如果 d=100,dₖ=100,r=20,那麼原本的計算時間是 d × dₖ = 10000,而重新投影的計算時間為 (d × r) + (dₖ× r) = 2000 + 2000 = 4000,提升了 60%

這不是新技巧;它在數學上類似於十年前 GoogleNet 提出的 *1×1 convolution1×1 卷積)理念。然而在 LLM 時代,其目的與 CNN 時代不同。在 CNN* 時代,我們可以把模型放在一張 GPU 上,瓶頸是推理時間。但在 LLM 時代,模型太大而無法放入單張 GPU,因此記憶體成為瓶頸。

在原本的 KV cache 中,我們為上千個頭和序列長度保留 dₖ 維的 KV 快取。但有了 MLA,我們只需要快取 r 維的壓縮版 KV,其中 r 遠小於 dₖ,如此節省的記憶體可以用於更多的 head 或更深層的結構。根據 *DeepSeek V2 論文,MLA 的計算成本與分成 2.25 組的 GQA 相近,但精度好於 GQA。根據他們實驗,MLA 可以讓原 67B 模型的 KV cache* 減少 93.3%。

來源: DeepSeek V2 paper

來源: DeepSeek V2 paper

啊然後咧?

第一次使用 ChatGPT 時,我們會對他擠牙膏的生成速度感到很不耐煩。但今天已經沒人抱怨他的回覆速度了,因為 LLM 背後的 MHA 計算已大幅提升;本文提到的 KV cache 只是其中一支分支。

KV cache 一直是 LLM 中重要的演算法,無論在系統工程還是演算法優化方面。最近的 MLA 才將 1×1 convolution 這一將近十年前的點子應用到 KV cache,這讓我相信過去十年的許多 CNN/DNN 的成果仍然能用在 LLM 上。雖然我們覺得 AI 與 LLM 炒作了很久,但是撇開大力出奇蹟的那些魔法,說不定從模型架構而言,我們還在起跑線的前端而已。

LLM 的雷之呼吸法終究會有幾型,還真說不清。


메타데이터
post_id
83f29c6b2a4b
slug
transformer-雷之呼吸第一型-kv-cache-83f29c6b2a4b
url
https://medium.com/@u9534056/transformer-%E9%9B%B7%E4%B9%8B%E5%91%BC%E5%90%B8%E7%AC%AC%E4%B8%80%E5%9E%8B-kv-cache-83f29c6b2a4b
canonical_url
https://medium.com/@u9534056/transformer-%E9%9B%B7%E4%B9%8B%E5%91%BC%E5%90%B8%E7%AC%AC%E4%B8%80%E5%9E%8B-kv-cache-83f29c6b2a4b
author_url
https://medium.com/@u9534056
status
ok
fetched_at
2026-06-12 07:40:50