Stable Diffusion JAX 與 Flax 整合
Hugging Face 已在 diffusers 函式庫中自版本 0.5.1 起整合了 Flax 支援,使 Stable Diffusion 能在 Google TPU 上高效運行。此整合讓使用者能利用 TPU 伺服器的平行運算能力(通常具備八個加速器),在產生單張圖像的時間內同時產生多張圖像。
使用 JAX 與 Flax 的高速 TPU 推理
透過使用 JAX 與 Flax,Stable Diffusion 在 TPU 上的推理已被最佳化,較標準 GPU 實作可顯著提升速度。在 TPU v2-8 上,首次編譯後的後續推理大約只需 7 秒。
主要技術優化包括:
- bfloat16 Precision:TPU 裝置支援
bfloat16,這是一種高效的半精度浮點類型,可在維持效能的同時減少記憶體開銷。 - JIT Compilation:將
jit=True傳遞給 Flax pipeline 後,JAX 會將模型編譯成高效的表示。首次執行需要編譯時間(在 TPU v2-8 上超過一分鐘),但之後的所有呼叫都會顯著加快。 - Stateless Models:由於 Flax 是函式式框架,模型是無狀態的,參數儲存在模型之外。
透過 SPMD 進行平行化
diffusers 的 Flax pipeline 採用單程式多資料 (Single-Program, Multiple-Data, SPMD) 平行化,以最大化 TPU 硬體使用率。這主要透過 jax.pmap 函式實現。
平行化的實作方式
jax.pmap 執行兩項關鍵功能:編譯程式碼(類似 jax.jit())以及確保編譯後的程式碼在所有可用裝置上平行執行。
為了平行執行,pipeline 依循以下步驟:
- Replication:模型參數透過
flax.jax_utils.replicate複製至所有裝置。 - Sharding:輸入資料(例如已分詞的提示 ID)使用
shard進行切分。例如,若有 8 個裝置,提示陣列會被分割,使每個裝置收到輸入的特定部分。 - PRNG Handling:為確保生成圖像的可重現性與多樣性,會建立隨機數生成器(RNG)並將其分割成多個生成器——每個裝置一個。
此架構使 pipeline 能同時產生八張不同的圖像(或同一圖像的八個副本),因為每個裝置會獨立處理一個 batch 項目。
模型存取與授權
Flax 版的 Stable Diffusion 權重可在 Hugging Face Hub 的 CompVis/stable-diffusion-v1-4 倉庫取得。存取需同意 CreativeML OpenRAIL-M 授權,其條款包括:
- 使用者不得故意使用模型產生或分享非法或有害內容。
- 使用者保留其生成輸出的權利,並對其使用負責。
- 允許商業使用與權重再分發,前提是必須將相同的使用限制與 CreativeML OpenRAIL-M 授權副本分享給所有使用者。