JOVANA
Explore Library Glossary Getting Started Three Levels Fields How it works Mission
Join the mission
All guides

塞得進記憶體的注意力:KV 快取的經濟學

解碼成本是記憶體問題,不是數學問題。認識 KV 快取,以及縮小它的 MQA/GQA、隱藏它的 FlashAttention-2、與限制它的滑動視窗。

為什麼解碼受限於記憶體

當 LLM 逐詞元生成時,每一步都對所有先前詞元重算注意力。為了不重做這份工,它把每一層的鍵與值存進 KV 快取(KV cache)。快取隨序列長度增長,且在整段生成期間絕不縮小,因此對長提示而言,瓶頸不是算術,而是把快取搬過記憶體。這是 LLM 服務的核心事實:吞吐量由位元組決定,而非浮點運算量。

標準的多頭注意力(multi-head attention)在每個位置、每個頭都保留各自的鍵與值向量。算上數十個頭與數十層,單一長對話的快取在記憶體流量上可以遠遠超過模型自身的權重。本篇的一切,都是對同一問題的不同答案:我們如何讓那份快取更便宜地儲存或搬移?

M_{\text{KV}} = 2\, b\, L\, n_{\text{layers}}\, n_{\text{heads}}\, d_{\text{head}}\, p

KV 快取的佔用隨序列長度、層數和注意力頭數線性增長——係數 2 來自鍵和值兩部分。

MQA 與 GQA:共用鍵與值

多查詢注意力(multi-query attention, MQA)採取了激進的一步:讓所有查詢頭共用單一的鍵/值頭。KV 快取縮小的倍數等於頭數——常是 8 到 64 倍——解碼速度大幅提升。代價是一些品質損失與偶發的訓練不穩定,因為所有那些表徵多樣性如今只能活在查詢裡。

標準多頭注意力為每個查詢頭都保留一個鍵/值頭;MQA 讓所有查詢頭共享一個鍵值頭,GQA 則每組共享一個。

多頭注意力示意圖:並行計算多個頭,拼接後再混合。

分組查詢注意力(grouped-query attention, GQA)是如今的標準折衷:把查詢頭分成少數幾組,每組各有自己共用的鍵/值頭。例如用 8 組,你幾乎能回收多頭注意力的全部品質,同時保住 MQA 大部分的快取節省。GQA 正是現代開源模型能在不爆記憶體的前提下服務長脈絡的原因——它是 Llama 2/3 與 Mistral 系列的預設。

# heads = 32 query heads; choose how many KV heads to keep
# MHA:  kv_heads = 32   (full cache, best quality)
# GQA:  kv_heads = 8    (4x smaller cache, ~MHA quality)
# MQA:  kv_heads = 1    (32x smaller cache, some quality cost)
kv_cache_bytes = kv_heads * head_dim * 2 * n_layers * seq_len * dtype_size
快取隨 KV 頭數而非查詢頭數縮放——這就是全部的槓桿所在。

FlashAttention-2:別再碰慢記憶體

MQA/GQA 縮小快取;FlashAttention-2則攻擊注意力本身如何搬資料。樸素的注意力會在慢的 GPU 記憶體中具現化完整的 N×N 分數矩陣,讀回來套用 softmax,再寫回去——對一個平方成長的矩陣做了三趟來回。第一卷的 FlashAttention 用分塊(tiling)與線上 softmax 把這些步驟融合,讓分數矩陣從不完全離開晶片上的快記憶體。結果是精確的注意力,搭配少得多的資料搬移。

\operatorname{Attention}(Q,K,V) = \operatorname{softmax}\!\left(\frac{QK^{\top}}{\sqrt{d_k}}\right)V

FlashAttention-2 分塊計算這一精確的 softmax 注意力,絕不在慢速顯存中生成完整的 N×N 分數矩陣。

*-2* 版是對同一想法的工程重思:它減少非矩陣乘的工作、除了批次與頭之外也沿序列維度平行化,並在 GPU warp 之間切分工作以削減共享記憶體流量。實務上它大致把原版的吞吐量翻倍,並把硬體使用率推向理論峰值。關鍵是它不改任何數學——輸出完全相同,所以它是幾乎總該開啟的免費收益。

滑動視窗注意力:給觸及範圍封頂

最便宜的詞元,是你從不去注意的那個。滑動視窗注意力(sliding-window attention)讓每個詞元只注意最近的 W 個詞元——一個固定的局部視窗——把注意力的成本從序列長度的平方降為線性,並把活躍的 KV 快取封頂在 W。資訊仍能傳得很遠,因為視窗會跨層疊加:經過 L 層後,一個詞元可受到大約 L 乘 W 個詞元之前的影響,就像深層卷積網路擴大其感受野那樣。