為什麼天真的解碼會一直做白工
要預測第 50 個詞元,注意力機制會拿這個新詞元跟每一個前面的詞元比對。若天真地做,每一步解碼都得從頭重跑整段序列——第 2 個詞元讀 1 個、第 1000 個詞元讀 999 個。總工作量隨長度的平方成長,而且你會把同樣的東西算了又算。
解法是記住。對每個過去的詞元,注意力需要兩個向量——一個鍵(key)與一個值(value)。它們一旦算出來就不會變,所以我們把它們存起來重複使用。這個儲存區就是KV 快取。
注意力作为查询匹配键和值并经过 softmax 加权的示意图。
快取如何讓迴圈變便宜
有了快取,預填充只算一次整段提示的鍵與值並存下來。之後每一步解碼只替那一個新詞元算鍵與值、接上去、再對整個快取做注意力。每步的工作量變成固定值而非持續成長——這正是讓長答案可行的那一個關鍵技巧。
代價:快取會吃掉你的 GPU
天下沒有白吃的午餐。快取要替每一個詞元、每一層、每一個注意力頭都存一組鍵與值——所以它的大小是*長度 × 層數 × 頭數 × 寬度 × 2*。KV 快取記憶體成本隨對話長度直線上升,在長脈絡或多使用者時,它甚至可能比模型權重本身還大。
KV 缓存随所有因素同时增长——序列长度、层数、注意力头数、宽度——再乘以二(键和值各一份)。
這就是為什麼大型語言模型的 GPU 記憶體是主宰服務的預算。一張卡要裝權重加上每個進行中請求的快取。當快取塞滿了卡,你就再也容不下一個使用者——GPU 的算力可能還閒著九成,卻已經完全沒有空間了。
# rough KV cache size for one sequence, fp16 (2 bytes/number) bytes = n_tokens * n_layers * 2 * n_kv_heads * head_dim * 2 # = length * layers * (K and V) * heads * width * fp16 # e.g. 4096 tokens, 32 layers, 8 KV heads, 128 dim -> ~0.5 GB per request
分頁:借用作業系統早就在用的把戲
早期的伺服器替每個請求預留一整塊連續記憶體,大小是照它可能給出的最長答案來算。但大多數答案都很短,所以那塊記憶體大半空著——浪費掉了。KV 快取分頁(以 PagedAttention 之名為人所知)的解法,是把快取切成固定大小的小區塊,就像作業系統把 RAM 切成分頁一樣。
區塊按需發放,請求一結束就立刻回收,因此幾乎沒有記憶體被閒置卡住。回報非常可觀:分頁通常能讓一張 GPU 同時容納數倍的並行請求,直接拉高每張卡能服務的使用者數。
快閃注意力:別把那張巨大矩陣蓋出來
注意力本身還藏著第二個記憶體陷阱。教科書做法會先在記憶體裡蓋出一張長度 × 長度的分數矩陣,再做 softmax——在長脈絡下那張矩陣可能有好幾 GB。快閃注意力(flash attention)從不真的把它蓋出來:它把注意力以小磚塊(tile)的方式串流,讓資料剛好塞進 GPU 微小的晶片內記憶體,一塊一塊地算出結果。
教科书式的做法会先在内存里构造 length×length 的 QK^T 分数矩阵再做 softmax;flash attention 算出相同结果,却从不把它整体存下来。
數學完全相同,改變的只是記憶體流量。藉由把運算留在快速的晶片內記憶體、而不是把巨大矩陣在卡的主記憶體之間來回搬,快閃注意力尤其加速了分數矩陣最大的預填充階段。分頁馴服了 KV 快取,快閃注意力馴服了注意力運算;兩者合起來決定了你能塞下多少。