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 快取透過在處理完初始提示後,為每一層儲存已計算的鍵與值,以減輕此低效。模型不再重新處理整個序列,而是遵循以下增量工作流程:
- 快取初始狀態: 首次傳遞後,為每層計算的 $K$ 與 $V$ 會被快取。
- 增量計算: 在隨後的生成步驟中,模型僅計算最新 token 的 $K$ 與 $V$。
- 快取更新: 新的 $K$ 與 $V$ 會被附加到現有快取中。
- 注意力計算: 當前 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 |
單階段生成迴圈 | 分為 prefill 與 decode 兩個階段 |
折衷與影響
KV 快取將每個 token 的推理複雜度從二次方降低至 $O(\text{seq len})$,使推理更快,且能在消費者硬體上執行大型模型。然而,此效率伴隨著折衝:需要更多記憶體來儲存快取,且增加了程式碼的複雜度。此外,它可能限制某些推理方案,例如需要更複雜快取管理的束搜索 (beam search)。
Sources
- OriginalKV Cache from scratch in nanoVLM