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 效能:
- 追蹤:隨著前向與反向傳遞的執行,中間表示(IR)圖會即時被追蹤。
- 截斷:當呼叫
xm.mark_step()(通常透過MpDeviceLoader間接呼叫)時,活躍圖會被切斷。 - 編譯:IR 圖被降低為 XLA 高階運算(HLO),編譯為 TPU 二進位檔並執行。
- 快取:為避免重新編譯的高成本,編譯後的 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 的限制。