注意力 KV 快取 分塊預填充 連續批次處理 結論
重點摘要:這篇文章從注意力機制和 KV 快取出發,透過優化吞吐量來推導出連續批次處理。
如果您曾使用過 Qwen、Claude 或任何其他 AI 聊天機器人,您可能注意到一件事:回應的第一個字出現需要一些時間,然後字詞會一個接一個地出現在您的螢幕上,(希望)以規律且快速的頻率出現。這是因為,歸根結底,所有 LLM 都只是花俏的下一個詞預測器。LLM 首先會處理您的整個提示,以產生一個新詞。然後,它會逐一添加詞,每次都會讀取之前的所有內容,直到它決定生成結束。
這個生成過程在計算上非常昂貴:每次生成一個詞都需要將輸入通過數十億個參數。為了使這些模型在實際應用中可行,特別是在同時服務多個用戶時,研究人員和工程師開發了一系列高效的推理技術。其中一項影響最大的優化是連續批次處理(continuous batching),它試圖透過並行處理多個對話並在它們完成時進行替換來最大化效能。
為了理解連續批次處理如何運作以及為何它在高負載服務場景中如此有效,我們將從 LLM 如何處理詞的基礎知識開始。
注意力機制是 LLM 運作的核心。語言模型透過將文本分解成我們稱為「詞」(tokens)的片段來處理文本。我們可以概念上將「詞」視為「單字」,但有時一個單字可能由幾個詞組成。對於每個詞序列,網路會計算下一個詞的預測。
網路中的許多操作都是逐詞的(token-wise):每個詞都是獨立處理的,給定詞的輸出僅取決於該詞的內容,而不取決於序列中的任何其他詞。像層歸一化或矩陣乘法等操作就是如此。然而,為了在句子中的單字之間建立聯繫,我們需要能夠讓詞相互影響的操作。
這就是注意力機制的作用。注意力層是不同詞相互作用的唯一地方。理解網路如何將詞連接起來,就意味著理解注意力。
讓我們看看在只有一個輸入提示的情況下,這實際上是如何運作的。
考慮初始提示「I am sure this project」,它被分解為 7 個詞:[<bos>, I, am, sure, this, pro, ject]。<bos>,或「序列開頭」(Beginning of Sequence),是我們在提示開頭添加的一個特殊詞,用來告知語言模型這裡開始了一個新的對話。
每個詞在網路內部都由一個長度為 d(隱藏維度)的向量表示。因此,這七個輸入詞形成一個形狀為 [1, n, d] 的張量 x。1 是序列數,或批次大小,在我們的情況下就是一。7 是序列長度,d 是隱藏維度,或每個詞表示的大小。接下來,我們將使用 n 而不是 7 作為序列長度。
輸入張量 x 然後通過三個矩陣進行投影:查詢投影 Wq、鍵投影 Wk 和值投影 Wv。這會產生三個張量 Q、K 和 V,它們的形狀都是 [1, n, A],其中 A 是注意力頭的維度。我們分別稱它們為查詢、鍵和值狀態。這在下圖的左側表示。
接下來,張量 Q 和 K 相乘以衡量詞之間的相似度,產生一個形狀為 [1, n, n] 的張量。這就是為什麼我們說注意力在序列長度上具有二次複雜度。計算 QKᵀ 需要 O(n²d) 的操作,因此成本是序列長度 n 的平方。這在上圖的右側表示。
然後,我們應用一個布林注意力遮罩(attention mask)到 QKᵀ 上,以控制哪些詞可以相互作用,如圖所示。在此圖中,注意力遮罩是因果遮罩(causal mask),意味著每個詞只能與其之前的詞相互作用。這符合因果必須先於結果的直覺,因此得名因果遮罩。注意力遮罩至關重要,因為它決定了網路中所有詞的互動。將所有注意力遮罩值設為 False,則整個網路中沒有任何詞會與其他詞互動。我們將在接下來的幾段中更仔細地檢視注意力遮罩。
最後,在應用注意力遮罩後,我們進行逐詞的 softmax(這與說逐列的 softmax 相同),並將結果乘以值投影 V,得到一個注意力頭的輸出,形狀為 [1, n, A]。我們在下圖中對整個過程進行視覺化總結。
我們將在這篇文章中使用大量的注意力視覺化,為了簡化,我們將稍微濃縮上圖。
為何重要:在連續批次處理中,Q、K 和 V 可以有不同數量的詞,因為,正如我們將看到的,我們將同時處理不同的階段(預填充和解碼)。為了使其更通用,讓我們假設 Q 的形狀為 [1, nQ, A],K 的形狀為 [1, nK, A],V 的形狀為 [1, nV, A]。
注意力分數 QKᵀ 的形狀則為 [1, nQ, nK],注意力遮罩也具有相同的形狀,因為它是逐點應用於分數的。
在應用注意力遮罩和逐列 softmax 後,我們乘以 V。由於我們將一個形狀為 [1, nQ, nK] 的矩陣乘以一個形狀為 [1, nV, A] 的矩陣,內部維度必須匹配:nK = nV。這意味著 V 和 K 的長度始終相同,因此我們可以透過僅顯示 K 來簡化我們的視覺化。如果這看起來很抽象,別擔心:圖表會讓它具體化。
此外,由於我們知道注意力遮罩是應用於 QKᵀ,我們知道它們具有相同的形狀。與其表示注意力分數,我們將在其位置表示注意力遮罩。最後,由於 Q、K 和 V 是 x 的直接投影,因此無需表示 x。這就得到了我們僅表示 Q、K 和注意力遮罩的簡化圖:
此表示法也強調了我們如何閱讀注意力遮罩。
我們逐列讀取遮罩,這與逐詞讀取相同:每一列對應一個詞的注意力計算。位置 (i, j) 的綠色方塊表示 True:詞 j 可以影響詞 i。白色方塊表示 False:不允許互動。
例如,查看詞「am」的第三列。「I」欄是綠色的,所以「I」影響了「am」的計算。「pro」欄是白色的,所以「pro」不影響「am」。這就是因果遮罩的作用:未來的詞無法影響過去的詞。
模型最後一層為每個輸入詞輸出一個詞預測。在我們處理單一提示的延續生成的情況下,我們只關心最後一個詞的下一個詞預測。在上面的圖中,最後一個詞是「ject」,相關的預測是「will」。
我們剛才描述的過程,即我們採用整個輸入序列,將其通過多個注意力層並計算下一個詞的分數,稱為預填充(prefill)。這是因為,正如我們稍後將看到的,我們執行的許多計算都可以被快取並重複使用——因此,我們正在預填充快取。由於使用了這個快取,序列生成可以以少得多的計算量進行,這個階段稱為解碼(decoding)。在解碼階段,生成一個新詞將比最初的全序列計算快得多。讓我們看看為什麼。
為了繼續生成,我們開始一個新的前向傳播,這通常看起來像這樣:
為了計算新詞的注意力分數,我們仍然需要先前詞的鍵和值投影。因此,我們需要重複舊詞(上圖中的灰色部分)與 Wk 和 Wv 的矩陣乘法,以檢索先前已計算過一次的結果。換句話說,我們浪費了計算資源。讓我們看看如何避免這種情況。
我們立即注意到,最後一個詞不會影響其他詞的注意力計算:
這符合因果遮罩的理念:由於「will」出現在所有先前詞之後,它不會改變它們的注意力計算。對於文本生成,因果注意力是迄今為止最常見的,因此我們將從現在開始專注於這種情況。請記住,非因果注意力方案也可以使用,尤其是在處理圖像時。考慮到我們只需要「will」這個詞的下一個詞預測,我們可以透過僅計算這個詞的輸出來簡化注意力機制。
此外,我們已經在先前的前向傳播中計算了「<bos>」…「ject」這些詞的 K 和 V 狀態:如果它們已被儲存,我們就不需要再次計算它們。這就是 KV 快取(KV cache):生成過程中創建的鍵和值狀態列表。它基本上允許透過避免重新計算鍵和值投影,將生成詞 n+1 的計算成本從 O(n²) 降低到 O(n),同時付出 O(n) 的記憶體成本。
在上圖中,只有白色部分的詞被計算:與計算 8 個詞的鍵和值相比,我們只計算了 1 個詞的鍵和值。您可以透過 KV 快取節省大量計算。您可以查看這篇文章以獲得更多 KV 快取的視覺化,或這篇文章以獲得實際實現範例。
讓我們更具體地說明快取大小,因為這是檢視我們模型中形狀的好機會。對於具有 L 個注意力層和 H 個具有頭維度 A 的注意力頭的模型,儲存一個詞所需的總快取大小為 2 \* L \* AH,其中 2 是為了同時考慮 K 和 V。例如,Llama-2-7B,L=32 層,H=32 頭,A=128,每層每詞需要 2 × 32 × 128 = 8,192 個值。使用 float16 精度,這需要 2AH × 2 位元組 = 16 KB 的記憶體。
KV 快取在我們想要生成下一個詞時很有用,這個階段我們稱之為解碼。但它在預填充階段也很有用,當我們處理初始提示並且有許多輸入詞時。特別是當初始提示很大,無法一次全部放入 GPU 記憶體時。
到目前為止,我們已經看了一個預填充的例子,其中 n=7 個詞,但在實際應用中,初始提示可能長得多。例如,在使用 Cursor 時,您可以將您的儲存庫添加到提示中作為上下文:這會顯著增加提示的大小。在這種情況下,儲存 n 個詞的激活所需的記憶體可能超過 GPU 的可用記憶體。因此,我們無法一次性完成預填充:我們必須將預填充分成幾個區塊。這稱為分塊預填充(chunked prefill),它將是實現高效推理所需的組件之一。
假設可用記憶體非常有限,我們每次前向傳播只能處理 m=4 個詞。如果我們有一個初始提示,其中 n=7 個詞,我們需要將其分成 ⌈n/m⌉ = 2 個區塊(將 7/4 = 1.75 向上取整為 2)。我們使用相同的 n 和 m 表示法來說明這個例子:
由於 KV 快取,我們可以做到這一點。我們在第一個預填充區塊中儲存 KV 狀態,在第二個預填充區塊中,我們將儲存的 KV 狀態附加到新的 KV 狀態之前。我們也相應地調整注意力遮罩。視覺化來看,這就像我們將非分塊預填充在中點分割。
關鍵洞察:快取的 KV 狀態讓我們能夠逐步處理提示,而不會丟失資訊。
雖然我們在這裡展示了一個將預填充分成 2 個區塊的例子,但分塊預填充可以用於以任何我們想要的方式分割預填充,靈活適應記憶體限制。
我們現在終於具備了理解連續批次處理所需的所有工具。
在我們之前的例子中,我們只考慮了批次大小為一的情況,即我們一次只為一個提示生成詞。在評估或模型服務的背景下,我們希望為大量提示生成詞。為了提高吞吐量(每秒生成的詞數),最佳做法是為一批提示並行生成詞。
為了將提示分組,樸素的方法是為輸入張量(詞序列和注意力遮罩)添加一個維度。然而,這對輸入的形狀有一個限制:我們需要所有提示都具有相同的長度,因為張量必須是矩形的。為了實現這一點,我們通常會在左側添加填充(padding),以便新的詞預測總是來自最右側的詞。我們也像下面這樣修改每個提示的注意力遮罩:
其中填充詞 <pad> 以橙色標示。然後,我們可以像以前一樣執行前向傳播,只是增加了批次大小的維度。這稱為批次生成(batched generation):對於相同長度的提示很有效,但當長度不同時則很浪費。這如下圖所示,透過 4 個生成步驟:一個預填充步驟(頂部)和 3 個解碼步驟(在每個「Forward pass」行下方)。
其中 <eos> 表示「序列結束」(End Of Sequence),這是一個特殊詞,用來指示模型已完成相應序列的生成。
批次生成的缺點是,如果一個提示比其他提示先完成生成(透過生成 <eos> 詞),則所有後續生成的詞都將無用。這會一直持續到批次中最長的請求完成為止。當然,我們可以從批次中移除已完成的提示並節省一些計算和記憶體,但這裡節省資源並不是目標:目標是吞吐量。
與其僅從批次中移除已完成的提示,不如用一個正在等待生成的提示來替換它。我們將稱之為動態調度(dynamic scheduling)或動態批次處理(dynamic batching)。動態調度對於在確保任何前向傳播生成的詞都相關的同時維持吞吐量非常有用。但由於我們將提示分組的方式,它有一個主要缺點:在交換提示時,我們需要大量的填充。這是因為新插入的提示需要經過預填充,而其他提示則一次解碼一個詞。因此,新插入提示的填充量幾乎與其詞數相同。
當批次大小增加且初始提示很長時,問題會變得更糟。填充成本會隨著批次大小和提示長度的增加而二次方增長。如果我們有一個包含 B 個提示的批次,這些提示都處於解碼階段,其中一個提示完成,動態地引入一個具有 n 個初始詞的提示,需要 (n-1)(B-1) 個填充詞。例如,對於 B=8 和 n=100,我們需要 99 × 7 = 693 個填充詞!
此外,像 CUDA graphs 或 torch.compile 這樣的實際優化需要靜態張量形狀。這迫使我們將所有提示填充到固定的最大長度,從而大大增加了填充浪費。
此時,我們的主要問題是填充,這是我們為了將句子分組而添加的維度的結果。因此,理想情況是完全擺脫這個維度,進行一次徹底的批次處理重新思考。如果我們這樣做,將提示分組的唯一方法是將它們串聯起來:
但我們不希望提示 0 的詞與提示 1 的詞互動!幸運的是,我們有一種方法可以控制詞如何相互作用:注意力遮罩。我們在下面展示了如何做到這一點:
雖然我們使用不同深淺的綠色來表示注意力遮罩的不同部分,但這仍然是一個布林遮罩,只有綠色表示 True,白色表示 False。這種將提示分組的方式稱為不規則批次處理(ragged batching)(因為序列長度是「不規則」或不均勻的),它提供了增加吞吐量的優勢,而無需引入填充詞的需求。
在上圖中,我們使用不規則批次處理將兩個完整的提示組合在一起,但我們可以根據記憶體允許的數量來批次處理任意數量的提示。唯一的限制是 m,即我們可以在一個批次中容納的詞數,其中 m 取決於 GPU 的可用記憶體。
不規則批次處理是連續批次處理的關鍵組件之一。為了最大化吞吐量,我們可以結合預填充和解碼序列,遵循類似以下的演算法:
動態調度是連續批次處理技術的最後一個組成部分:我們在完成的提示一完成就將它們從批次中移除,並用對應於進入請求的新分塊提示替換它們。
這種不規則批次處理和動態調度的結合稱為連續批次處理,它是現代 LLM 服務系統的動力所在。
連續批次處理結合了三項關鍵技術來最大化 LLM 服務的吞吐量:
透過移除批次維度並使用注意力遮罩來控制詞的互動,連續批次處理允許在同一個批次中混合預填充和解碼階段,從而顯著提高服務多個請求的效率。這就是為什麼像 ChatGPT 這樣的服務能夠高效處理數千個並發用戶。
在我們這個系列的下一篇文章中,我們將透過引入非同步批次處理(asynchronous batching)來使連續批次處理更加高效。在此處閱讀!如果您想深入了解其他連續批次處理的主題,請在評論中告訴我們!