Hugging Face 將 LLM.int8() 8 位矩陣乘法整合至 Transformers 與 Accelerate

TL;DR

Hugging Face 宣佈,8 位元 LLM.int8() 量化方法現在已完整整合至 transformersaccelerate 函式庫,讓諸如 BLOOM‑176B 等巨型模型的推論可使用約一半的記憶體佔用,且在精度上沒有可測量的損失。


為何 8 位元量化對大型語言模型重要

大型語言模型(LLM)現在已超過數千億參數(例如 PaLM 540B、OPT 176B、BLOOM 176B)。以全精度 FP32 儲存模型需要每個權重 4 位元組,導致記憶體需求達到數百 GB,遠超大多數 GPU 的容量。將精度降低至半精度(FP16/BF16)可將記憶體減半,但像 BLOOM 176B 仍需約 350 GB。使用 8 位元整數(INT8)量化則可再減少 2 倍,但傳統的簡單量化往往會降低精度,尤其是對於超過約 6 B 參數的模型。

LLM.int8() 的核心概念:零退化矩陣乘法

LLM.int8() 透過將 異常值(outlier)單獨處理,克服了精度下降的問題:

  1. 異常值提取 – 針對隱藏狀態矩陣的每一列,識別出幅度超過閾值(≈6)的值。
  2. 混合精度矩陣乘法 – 異常值以 FP16 進行乘法,其餘大部分矩陣則量化為 INT8,並使用向量級(對於激活為行級,對於權重為列級)量化進行乘法。
  3. 去量化與聚合 – 將 INT8 結果去量化回 FP16,並與異常值的 FP16 結果相加,產生最終的 FP16 輸出。

此三步驟流程在保留原始 FP16/BF16 模型精確推論品質的同時,將記憶體使用量降低至四分之一。

量化機制:零點 vs. 絕對最大值

  • 零點量化 將浮點範圍(例如 [-1, 1])縮放至 INT8 範圍 [-127, 127],並對每個值進行四捨五入。逆向縮放可恢復原始值的近似。
  • 絕對最大值量化 先將每個張量除以其絕對最大值,再乘以 127,最後四捨五入。對於向量 [1.2, ‑0.5, ‑4.3, …, 5.4],縮放因子為 127/5.4 ≈ 23.5,產生介於 [-127, 127] 的整數值。

兩種方案皆可行於行級或列級,這對於大規模的精確矩陣乘法至關重要。

退化的實證證據

在 OPT‑175B 與 BLOOM‑176B 上使用 lm‑eval‑harness 的基準測試顯示,INT8 與 FP16/BF16 分數之絕對差異在所有任務中皆低於標準誤差(例如 HellaSwag 準確率 0.7849 vs. 0.7849,Lambada 困惑度 3.0142 vs. 3.0152)。在一個案例(BLOOM‑176B 在 Lambada)中,INT8 模型的表現略好於 FP16。論文 LLM.int8(): 8‑bit Matrix Multiplication for Transformers at Scale 提供了完整的評估。

效能取捨

記憶體節省會伴隨對最大模型的輕微速度下降:BLOOM‑176B 在 INT8 模式下比 FP16 慢 15 %–23 %。較小的模型(例如 T5‑3B、T5‑11B)最初的速度下降較大,但近期的最佳化將每個 token 的延遲從 312 ms 降至 173 ms(T5‑3B),以及從 45 ms 降至 25 ms(T5‑11B)。未來的版本將進一步縮小這個差距。

模型 精度 GPU 數量 Tokens / ms(批次 1)
BLOOM‑176B BF16 8 × A100 80GB 239
BLOOM‑176B INT8 4 × A100 80GB 282
T5‑11B FP16 2 × T4 15GB 11.7
T5‑11B INT8 1 × T4 15GB 43.5

整合至 transformers

關鍵組件是 bitsandbytes.nn.Linear8bitLt,它是 torch.nn.Linear 的即插即用替代品。以下是一個最小化的轉換工作流程:

import torch, bitsandbytes as bnb
from bnb.nn import Linear8bitLt

# Define a FP16 model and save its weights
fp16 = torch.nn.Sequential(torch.nn.Linear(64, 64), torch.nn.Linear(64, 64))
torch.save(fp16.state_dict(), "model.pt")

# Build an INT8 version
int8 = torch.nn.Sequential(
    Linear8bitLt(64, 64, has_fp16_weights=False),
    Linear8bitLt(64, 64, has_fp16_weights=False),
)
int8.load_state_dict(torch.load("model.pt"))
int8 = int8.to(0)   # quantization occurs on GPU

在呼叫 .to 之後,權重會以範圍 [-127, 127] 的 int8 張量儲存。原始的 FP16 值可透過 (weight.CB * weight.SCB) / 127 復原。

利用 accelerate 進行零記憶體模型建構

accelerate.init_empty_weights() 會在 meta 裝置上建立模型,且不分配任何記憶體。此整合會修補 accelerate,使參數在從 meta 裝置移除時仍保留其自訂類別(Int8Params)。遞迴輔助函式會將每個 nn.Linear 替換為 Linear8bitLt,同時保留如 lm_head 等應保持全精度的模組:

from accelerate import init_empty_weights
import torch.nn as nn, bitsandbytes as bnb

def replace_8bit_linear(model, threshold=6.0, exclude="lm_head"):
    for name, module in model.named_children():
        if list(module.children()):
            replace_8bit_linear(module, threshold, exclude)
        if isinstance(module, nn.Linear) and name != exclude:
            with init_empty_weights():
                model._modules[name] = bnb.nn.Linear8bitLt(
                    module.in_features,
                    module.out_features,
                    module.bias is not None,
                    has_fp16_weights=False,
                    threshold=threshold,
                )
    return model

兩個針對 accelerate 的 PR 確保對每個 INT8 張量僅呼叫一次 set_module_tensor_to_device,以避免雙重量化的錯誤。

硬體與安裝需求

  • GPU 支援 – 需要 INT8 張量核心(NVIDIA Turing、Ampere、RTX 20/30、A40‑A100、T4)。CPU 與較舊的 Kepler GPU 不具備原生支援。
  • 安裝 – 使用 Python ≥ 3.8:
pip install accelerate bitsandbytes
pip install git+https://github.com/huggingface/transformers.git

示範

Google Colab 筆記本展示了使用 INT8 執行 T5‑11B(原本在 FP32 下需 42 GB)僅佔 11 GB,及一個可在單一 T4 上順利運行的 BLOOM‑3B 示範。

未來工作與限制

  • 小模型的速度 – 持續的工作旨在使 ≤6 B 模型的 INT8 延遲與 FP16 相當。
  • Kepler GPU 支援 – 計畫為缺乏原生 INT8 張量核心的 GPU(例如 GTX 1080)加入獨立的軟體堆疊。
  • State‑dict 持久化 – 目前的 INT8 檢查點省略了量化統計資訊(CBSCB),導致無法直接從 Hub 載入;加入此中繼資料是首要任務。
  • CPU 執行 – CPU 上沒有 8 位元張量核心;未來的軟體路徑可能擴大可及性。
  • 超越文字 – 將此技術擴展至大型視覺、音訊與多模態模型仍是未解的研究方向。

致謝:Younes B.、Tim Dettmers,以及原始部落格文章中列出的貢獻者。

Sources