在生產環境中優化 LLMs:精度、注意力與架構
在生產環境中部署大型語言模型 (LLMs) 需要克服兩個主要瓶頸:數十億參數帶來的巨顯存需求以及長輸入序列相關的二次記憶體增長。為了解決這些問題,Hugging Face 建議結合低精度量化、優化的注意力算法以及策略性的架構選擇。
透過降低精度減少記憶體佔用
降低數值精度可減少載入模型權重所需的顯存,使得更大的模型能夠在較小或更易於存取的硬體上運行。
各精度的顯存需求
載入模型權重是短文本輸入(低於 1024 個 token)的主要記憶體成本。顯存需求的經驗法則如下:
- float32:對於具有 X 十億參數的模型,大約需要 4 * X GB 的顯存。
- bfloat16/float16:對於具有 X 十億參數的模型,大約需要 2 * X GB 的顯存。
例如,Llama-2-70b 在 bfloat16 下需要約 140 GB 的顯存,超過單張 A100(80 GB)的容量,因而需要使用張量或管線平行。
量化(8 位元和 4 位元)
量化將精度進一步降低至 8 位元或 4 位元,顯著降低記憶體使用量,對文字生成的準確度影響極小。這是因為文字生成依賴於下一個 token 的 logits 的相對分布,而非精確值。
- 8 位元量化:顯著降低顯存使用量(例如,OctoCoder 的峰值記憶體從約 32 GB 降至約 15 GB)。由於在計算過程中需要動態反量化,可能會導致推理時略微變慢。
- 4 位元量化:進一步降低顯存(例如,OctoCoder 降至約 9.5 GB),使得模型能夠在消費級 GPU(如 RTX 3090)上運行。然而,與 8 位元量化相比,它可能導致更明顯的準確度下降和較慢的推理。
使用 Flash Attention 加速推理
標準自注意力相對於序列長度 ($N$) 具有二次計算和記憶體複雜度,這使得在長上下文(例如 16,000+ 個 token)時變得極其昂貴。
Flash Attention 演算法
Flash Attention 透過將計算分割成較小的區塊並多次迭代 softmax 步驟來優化注意力機制。它避免了大型 $QK^T$ 矩陣的產生,使得記憶體成本隨著 $N$ 的增加呈線性而非二次增長。
效能提升
儘管 Flash Attention 由於重新計算 softmax 正規化統計而需要更多的 FLOPs,但在實際應用中它更快,因為它減少了對慢速高頻寬記憶體(顯存)的存取,並最大化利用快速的片上 SRAM。它產生的輸出與預設自注意力算法在數值上完全相同。
長上下文與聊天的架構優化
訓練期間所做的架構選擇決定了模型處理長序列和多輪對話的效率。兩個關鍵領域是位置嵌入和鍵值(KV)快取。
相對位置嵌入
絕對位置嵌入(正弦波或學習得到的)在處理長文本時通常表現不佳,且難以超越其訓練長度進行外推。相對位置嵌入則更為有效:
- 旋轉位置嵌入 (RoPE):通過旋轉 query-key 對來編碼位置。它被用於 Falcon、Llama 和 PaLM。
- ALiBi:將預定義值的負整數縮放後加到 $QK^T$ 矩陣中。它被用於 MPT 和 BLOOM,通常相比 RoPE 更有效地外推到更長的序列。
優化鍵值(KV)快取
自回歸生成使用 KV 快取來存儲所有先前 token 的 key-value 向量,從而避免在每一步重新計算它們。這將 $QK^T$ 的計算轉換為向量-矩陣乘法 ($\text{query} \times \text{KV cache}$),顯著提升速度。
然而,KV 快取可能成為記憶體瓶頸。兩種架構可降低此開銷:
- 多查詢注意力 (MQA):使用單個鍵值投影頭,在所有注意力頭之間共享。這大幅減少了快取大小(例如,對於 OctoCoder 中的 16,000 個 token 序列,從 15 GB 減少到低於 400 MB),並減少記憶體頻寬瓶頸。用於 Falcon、PaLM、MPT 和 BLOOM。
- 分組查詢注意力 (GQA):介於 MQA 和標準多頭注意力之間的折衷方案。它使用少量的 KV 投影頭(例如 2、4 或 8),以保持比 MQA 更大的模型容量,同時保留其大部分效率。用於 Llama-2。