nanoVLM 中的 KV 快取實作

nanoVLM 中的 KV 快取實作

Hugging Face 從頭在 nanoVLM 倉庫中實作了 KV 快取,使其視覺語言模型的生成速度提升了 38%。

nanoVLM 中的 KV 快取實作

Hugging Face 從頭在 nanoVLM 中實作了 KV(鍵值)快取,這是一個用於訓練視覺語言模型的精簡 PyTorch 程式碼庫。此優化透過在自回歸推理過程中消除冗餘計算,使生成速度提升了 38%

自回歸生成中的計算冗餘

自回歸語言模型一次產生一個 token。在未使用快取的標準 transformer 實作中,模型必須處理整個序列——包括所有先前產生的 token——才能預測下一個 token。

由於 transformer 在內部是平行運算,每一次新 token 的預測都需要對所有層進行完整的前向傳播。這會導致相對於序列長度的記憶體與計算需求呈二次方增長。具體而言,模型在每一步都會重新計算所有先前 token 的鍵 (K) 與值 (V) 張量,即使這些 token 及其對應的投影並未改變。

KV 快取如何優化推理

KV 快取透過在處理完初始提示後,為每一層儲存已計算的鍵與值,以減輕此低效。模型不再重新處理整個序列,而是遵循以下增量工作流程:

  1. 快取初始狀態: 首次傳遞後,為每層計算的 $K$ 與 $V$ 會被快取。
  2. 增量計算: 在隨後的生成步驟中,模型僅計算最新 token 的 $K$ 與 $V$。
  3. 快取更新: 新的 $K$ 與 $V$ 會被附加到現有快取中。
  4. 注意力計算: 當前 token 的查詢 ($Q$) 會與快取的 $K$ 與 $V$ 結合,以產生輸出。

實務上,這個快取以每層字典的形式維護,包含形狀為 (batch_size, num_heads, seq_len_cached, head_dim) 的「key」與「value」張量。

nanoVLM 中的技術實作

nanoVLM 的實作涉及對三個主要元件的修改,以從完整序列的重新計算轉變為增量更新系統。

1. 注意力區塊更新

LanguageModelGroupedAttention 類別中,forward 函式被修改以接受 block_kv_cache。若快取已存在(表示模型不在預填階段),模型會為當前 token 計算 $K_{new}$ 與 $V_{new}$,並將它們與快取的張量串接。若快取不存在,則對提示執行初始計算。

2. 層級快取追蹤

LanguageModel 類別現在實作層級快取追蹤。它使用 start_pos 參數,以確保旋轉位置編碼與當前生成索引正確對齊,使模型能知道新生成 token 相對於序列的絕對位置。

3. 生成迴圈的分支

VisionLanguageModel 中的 generate() 方法被分為兩個不同階段:

  • 預填階段: 模型編碼完整的輸入提示,並為所有層構建初始 KV 快取。
  • 解碼階段: 模型逐步產生 token,使用快取的鍵和值,以避免重新處理提示與先前生成的 token。

架構變更摘要

模組 原始行為 新行為
LanguageModelGroupedAttention.forward 在每一步重新計算 $Q$, $K$, $V$ 使用並更新 KV 快取
LanguageModel.forward 未保留先前狀態 追蹤每層 KV 快取,處理 start_pos
VisionLanguageModel.generate 單階段生成迴圈 分為 prefilldecode 兩個階段

折衷與影響

KV 快取將每個 token 的推理複雜度從二次方降低至 $O(\text{seq len})$,使推理更快,且能在消費者硬體上執行大型模型。然而,此效率伴隨著折衝:需要更多記憶體來儲存快取,且增加了程式碼的複雜度。此外,它可能限制某些推理方案,例如需要更複雜快取管理的束搜索 (beam search)。

Sources