微調 LLM 至 1.58 位元:使用 BitNet 的極端量化

Hugging Face 開發了一種方法,利用 BitNet 架構將現有的大型語言模型(LLM)微調至 1.58 位元精度。此方法使模型僅以三個值(-1、0、1)表示參數,極大降低計算與能源成本,且不需要通常用於從頭預訓練 1 位元模型的龐大預算。

BitNet 架構與 1.58 位元量化

BitNet 用 BitLinear 層取代多頭注意力(Multi-Head Attention)與前饋網路(Feed-Forward Networks)中的標準 Linear 層。這些層對權重使用三元精度,對激活使用 8 位元精度。

計算範式

與依賴 FP16 加法與乘法的標準 LLM(例如 Llama)不同,BitNet b1.58 在矩陣乘法中使用 INT8 加法。此計算方式的轉變在理論上可將矩陣乘法的能源消耗降低至 Llama 基線的 71.4 倍。

使用 Straight-Through Estimator(STE)進行訓練

由於用於三元量化的 round() 函式不可微分,BitNet 採用 Straight Through Estimator(STE)。STE 將捨入操作的梯度近似為 1,使梯度能如同通過恆等函式般流過該操作,從而支援標準的基於梯度的優化。

量化機制

  • 權重: 使用對稱的 per‑tensor 量化。尺度為權重矩陣絕對值平均值的倒數。權重先被縮放、捨入、限制在 -1 到 1 之間,最後再重新縮放。
  • 激活: 以 absmax per‑token 量化方式量化至 8 位元精度,將數值縮放至 [-128, 127] 範圍。激活量化前會先套用層正規化(Layer Normalization,LN),以維持輸出變異。

微調現有模型至 1.58 位元

Hugging Face 成功將 Llama 3 8B 模型微調至 1.58 位元精度。初始實驗發現,若突然加入 BitLinear 層會導致模型幾乎失去所有預訓練資訊,從而出現損失急升。

動態暖身量化

為防止先前知識的流失,Hugging Face 實作了一個動態 $\lambda$(lambda)值,以逐步引入量化:

$$\lambda = \min\left(\frac{\text{training_step}}{\text{total_training_steps}}, 1\right)$$

透過以 $\lambda$ 縮放原始值與量化值之差,模型會從全精度($\lambda=0$)過渡到全量化($\lambda=1$)。此線性排程器提升了收斂效果,使 TinyStories 資料集的困惑度約為 4。

擴展與泛化

為確保模型保有通用知識且不會對小型資料集過度擬合,團隊將訓練規模擴展至 FineWeb-edu 資料集。以 1e-4 的學習率、每批 200 萬 token,總計 100 億 token 訓練,模型在 WikiText 上的困惑度達到 12.2。

進一步擴展至 1000 億 token 的實驗顯示,雖然模型在某些指標上與原始 Llama 3 8B 相近,整體仍略遜於全精度基線。

效能基準與結果

使用 1.58 位元架構微調的模型已於 HF1BitLLM 組織下發布。

主要發現

  • 具競爭力的效能: 在 100 億 token 微調後,1.58 位元 Llama 3 8B 模型的表現超過了在 1000 億 token 上訓練的 BitNet 7B 模型以及在 1.26 兆 token 上蒸餾的 FBI LLM。
  • MMLU 基準: 開發的 8B 模型在 MMLU 基準上超過了 Llama 1 7B 模型。
  • 模型大小: 將權重打包成 int8 張量,使參數量從 80 億降至 28 億。

推論最佳化與自訂核心

為實現 1.58 位元權重的速度與記憶體優勢,Hugging Face 實作了自訂的 CUDA 與 Triton 核心,以在矩陣乘法過程中即時解壓權重。

瓦片化矩陣乘法

為克服記憶體頻寬瓶頸與冗餘資料存取,團隊使用 tiling(瓦片化)技術。此方法將矩陣切分為可容納於 GPU 快速共享記憶體的較小子矩陣(瓦片),降低慢速全域記憶體存取的頻率。

核心效能測試

  • Triton vs. Torch: 自訂的 Triton 核心的效能大致等同於使用 BF16 精度的 @torch.compile
  • BitBlas: 團隊發現 BitBlas(一個混合精度軟體庫)在低精度下的效能優於自訂 Triton 核心與 Torch 的 matmul 函式,儘管因核心編譯導致較高的初始載入時間。

與 Transformers 的整合

整合透過 transformers 套件中的全新「bitnet」量化方法實現。標準 Linear 層被專用的 BitLinear 層取代。API 保持不變,使用者可透過 AutoModelForCausalLM.from_pretrained 載入模型。

Sources