為什麼解碼受限於記憶體
當 LLM 逐詞元生成時,每一步都對所有先前詞元重算注意力。為了不重做這份工,它把每一層的鍵與值存進 KV 快取(KV cache)。快取隨序列長度增長,且在整段生成期間絕不縮小,因此對長提示而言,瓶頸不是算術,而是把快取搬過記憶體。這是 LLM 服務的核心事實:吞吐量由位元組決定,而非浮點運算量。
標準的多頭注意力(multi-head attention)在每個位置、每個頭都保留各自的鍵與值向量。算上數十個頭與數十層,單一長對話的快取在記憶體流量上可以遠遠超過模型自身的權重。本篇的一切,都是對同一問題的不同答案:我們如何讓那份快取更便宜地儲存或搬移?
KV 快取的佔用隨序列長度、層數和注意力頭數線性增長——係數 2 來自鍵和值兩部分。
MQA 與 GQA:共用鍵與值
多查詢注意力(multi-query attention, MQA)採取了激進的一步:讓所有查詢頭共用單一的鍵/值頭。KV 快取縮小的倍數等於頭數——常是 8 到 64 倍——解碼速度大幅提升。代價是一些品質損失與偶發的訓練不穩定,因為所有那些表徵多樣性如今只能活在查詢裡。
多頭注意力示意圖:並行計算多個頭,拼接後再混合。
分組查詢注意力(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
FlashAttention-2:別再碰慢記憶體
MQA/GQA 縮小快取;FlashAttention-2則攻擊注意力本身如何搬資料。樸素的注意力會在慢的 GPU 記憶體中具現化完整的 N×N 分數矩陣,讀回來套用 softmax,再寫回去——對一個平方成長的矩陣做了三趟來回。第一卷的 FlashAttention 用分塊(tiling)與線上 softmax 把這些步驟融合,讓分數矩陣從不完全離開晶片上的快記憶體。結果是精確的注意力,搭配少得多的資料搬移。
FlashAttention-2 分塊計算這一精確的 softmax 注意力,絕不在慢速顯存中生成完整的 N×N 分數矩陣。
*-2* 版是對同一想法的工程重思:它減少非矩陣乘的工作、除了批次與頭之外也沿序列維度平行化,並在 GPU warp 之間切分工作以削減共享記憶體流量。實務上它大致把原版的吞吐量翻倍,並把硬體使用率推向理論峰值。關鍵是它不改任何數學——輸出完全相同,所以它是幾乎總該開啟的免費收益。
滑動視窗注意力:給觸及範圍封頂
最便宜的詞元,是你從不去注意的那個。滑動視窗注意力(sliding-window attention)讓每個詞元只注意最近的 W 個詞元——一個固定的局部視窗——把注意力的成本從序列長度的平方降為線性,並把活躍的 KV 快取封頂在 W。資訊仍能傳得很遠,因為視窗會跨層疊加:經過 L 層後,一個詞元可受到大約 L 乘 W 個詞元之前的影響,就像深層卷積網路擴大其感受野那樣。