Hugging Face PyTorch / XLA TPU 整合

Hugging Face 已將 PyTorch / XLA 整合,以允許使用者在保持標準 Hugging Face Trainer 介面的同時,於 Cloud TPU 上訓練與擴展 transformer 模型。此整合利用 PyTorch / XLA 庫將 PyTorch 框架與 XLA(加速線性代數)設備(包括 Cloud TPU)連接起來。

PyTorch / XLA 技術實作

此整合將 xla 裝置類型引入 PyTorch,使得張量能夠在 TPU 硬體上建立與管理。Hugging Face Trainer 模組利用 TrainingArguments 資料類別,在 is_torch_tpu_available() 為 true 時自動偵測並返回 TPU 裝置。

梯度合併與優化器步驟

由於 Cloud TPU 裝置通常由多個核心組成(例如,單個裝置可能擁有 8 個核心),必須在資料平行副本之間交換梯度。此整合使用 xm.optimizer_step(optimizer) 來處理梯度合併及後續的優化器步驟,確保 TPU 核心間的同步。

輸入管線

為防止主機 CPU 與 TPU 加速器在互相等待時閒置,PyTorch / XLA 實作了一個輸入管線。透過使用 pl.MpDeviceLoader,系統能夠在步驟 $n$ 仍在執行時重疊步驟 $n+1$ 的追蹤,從而優化模型的資料輸入。

檢查點管理

為確保可攜性並避免裝置特定載入問題,張量在檢查點之前會被移至 CPU。使用 xm.save() API 確保只有單一進程(主序數)寫入儲存位置,以防止在多進程環境中發生檔案損毀。

PyTorch / XLA 的運作方式

惰性張量執行

與 CPU 與 CUDA 張量不同,XLA 張量是惰性的。它們會在結果需要之前將操作記錄在圖中。這種延遲執行使得 XLA 編譯器能夠將多個獨立操作融合為單一最佳化操作。

Trace-Compile-Execute 週期

PyTorch / XLA 遵循特定的執行流程以優化 TPU 效能:

  1. 追蹤:隨著前向與反向傳遞的執行,中間表示(IR)圖會即時被追蹤。
  2. 截斷:當呼叫 xm.mark_step()(通常透過 MpDeviceLoader 間接呼叫)時,活躍圖會被切斷。
  3. 編譯:IR 圖被降低為 XLA 高階運算(HLO),編譯為 TPU 二進位檔並執行。
  4. 快取:為避免重新編譯的高成本,編譯後的 TPU 二進位檔會依據 HLO 圖的唯一雜湊值存放在快取中。

為最大化快取命中並最小化編譯開銷,建議保持張量形狀靜態。Hugging Face 模型通常透過適當填充輸入標記來維持靜態形狀。

效能基準

在 v3-8 Cloud TPU 系統(4 顆 TPU v3 晶片)上使用 WikiText103 資料集訓練 bert-large-uncased 會得到以下結果:

名稱 全域批次大小 精度 訓練時間(分鐘)
bert-large-uncased 64 FP32 178.4
bert-large-uncased 128 BF16 106.4

這些基準是使用 n1-standard-96 CPU 配置進行的,以確保工作負載不會受到主機 CPU 的限制。

Sources