在 Cloud TPU v5e 上使用 JAX 加速 Stable Diffusion XL 推理
Hugging Face 已將 JAX 支援整合到 Diffusers 函式庫中,以在 Cloud TPU v5e 上實現 Stable Diffusion XL (SDXL) 的高效能、具成本效益的推理。此整合透過利用 JAX 的即時 (JIT) 編譯和 XLA 驅動的平行ism,解決了 SDXL 的計算挑戰,其 UNet 大小約為前身的三倍。
透過 JAX 和 TPU v5e 的技術優化
在 Cloud TPU v5e 上提供 SDXL 透過兩個主要的軟體與硬體機制實現高效能:JIT 編譯和 SPMD 平行ism。
靜態形狀的 JIT 編譯
JAX 利用即時 (JIT) 編譯在初始執行期間追蹤程式碼,並為後續呼叫生成最佳化的 TPU 二進位檔。此過程需要靜態的輸入、中間和輸出形狀。SDXL 與 JIT 編譯高度相容,原因如下:
- 固定輸出形狀: 圖像生成通常使用固定數量的圖像和一致的尺寸。
- 固定形狀嵌入: Stable Diffusion 和 SDXL 使用固定形狀的嵌入向量(帶填充)來處理文本提示。
雖然初始編譯需要數分鐘(在提供的範例中約為三分鐘),但後續的推理呼叫將顯著加速。
XLA 平行ism 和吞吐量
JAX 的 pmap 能啟用單程式多資料 (SPMD) 執行,使工作負載能夠跨多個 XLA 裝置進行縮放。這使得圖像生成可以線性縮放:例如,擁有 8 個晶片的 TPU 在單個晶片生成一張圖像所需的時間內,可以生成 8 張圖像。Cloud TPU v5e 實例提供多種配置(從 1 到 256 個晶片),透過超高速 ICI 連結相連,讓使用者能根據特定的吞吐量需求進行擴展。
在 JAX 中的實作管線
使用 JAX 運行 SDXL 推理涉及一種功能方法,其中模型參數與管線分開處理。主要實作步驟包括:
- 模型載入: 使用
FlaxStableDiffusionXLPipeline.from_pretrained載入基礎 SDXL 1.0 模型。 - 精度管理: 將模型參數轉換為
bfloat16以減少記憶體使用並提升速度,同時將排程器狀態保持為float32,以防止導致低品質或黑色圖像的精度誤差。 - 輸入準備: 使用
prepare_inputs確保提示在各次呼叫中具有一致的維度,這是 JIT 編譯所必需的。 - 設備複寫: 在可用的 TPU 晶片間複寫參數和輸入(例如,對 TPU v5e-4 使用
replicate),並為每個晶片分配唯一的隨機種子,以確保圖像輸出的多樣性。 - 執行: 使用
jit=True呼叫管線以觸發 XLA 編譯過程。
效能基準測試
在使用 Euler 離散排程器進行 20 步的 SDXL 1.0 基礎版基準測試顯示,TPU v5e 在成本效益方面優於 TPU v4。
| 硬體 | 批次大小 | 延遲 | 每美元性能 |
|---|---|---|---|
| TPU v5e-4 (JAX) | 4 | 2.33s | 21.46 |
| TPU v5e-4 (JAX) | 8 | 4.99s | 20.04 |
| TPU v4-8 (JAX) | 4 | 2.16s | 9.16s |
| TPU v4-8 (JAX) | 8 | 4.17s | 8.98 |
TPU v5e 每美元的性能最高可比 TPU v4 高出 2.4 倍。性能是透過計算吞吐量(批次大小除以每個晶片的延遲),然後將該數值除以硬體的標價來測量的。
部署架構
目前的實作使用一個負載平衡伺服器,將使用者請求隨機路由至運行在預先分配的 Cloud TPU v5e-4 實例上的後端伺服器。每個實例在約 4 秒內生成四張 1024×1024 的圖像(包括前端處理與通訊),實際生成時間約為 2.3 秒。