返回首頁

AI 基礎設施/分散式運算

Ray 將 TPU slice 排程下沉至 Serve、Data 與 Train,統一多主機 AI 工作流

Google 公布 Ray 在 TPU 上的完整 AI 程式庫路徑,以 topology 欄位確保模型工作者落在同一 ICI slice。新版並加入 JAX 原生資料批次、JaxTrainer、官方容器與 TPU 監控,但目前主要驗證環境仍是 GKE。

Internet Archive Book Images · No restrictions · Image source
zh-Hant

Ray 2.55 的 TPU 支援不只讓一般 task 與 actor 看見加速器;Google 7 月 24 日公布的第二階段整合,把 TPU 的拓撲限制帶進 Ray Serve、Ray Data 與 Ray Train。TPU 晶片以固定 slice 透過 ICI 高速互連,多主機張量平行模型必須完整落在同一 slice。若排程器把逐晶片資源分散到兩個 slice,第一次 collective 便可能永遠等不到,而部署只會停在 `DEPLOYING`。

Ray Serve 現在可在 `accelerator_config` 指定例如 `topology: "4x4"`,由 replica 建立 slice placement group,對整個拓撲執行 gang scheduling;模型服務仍可使用 vLLM 的負載平衡、擴縮與多模型組合。Ray Data 新增 `iter_jax_batches()`,直接產生已轉為 JAX array 並完成裝置分片的批次,減少 NumPy 到 JAX 的主機端複製,也明確處理最後一個不規則批次的捨棄、填補或報錯策略。

訓練端的 `JaxTrainer` 則把 slice 形狀、每主機工作者、跨 slice 協調、checkpoint 與容錯包進 `ScalingConfig`。Google 同步提供帶有 JAX、Flax、Optax、Orbax 與分析工具的 `rayproject/ray:*-tpu` 映像;Ray Dashboard 也能呈現 TPU 利用率與記憶體。公開範例以單一 TPU v6e `2x4` slice、Qwen3-4B、LoRA/DPO 串起資料準備、微調、批次推論與服務。

這降低了既有 Ray 團隊跨到 TPU 的排程與環境組裝成本,也讓同一 Python 抽象可跨 GPU、TPU 工作流。不過「只加 topology」並不消除 TPU 的靜態形狀、編譯時間與跨 slice 網路成本;官方路徑目前又高度依賴 GKE、KubeRay 和 Google 提供的映像。工程團隊接下來應驗證大型模型的實際吞吐、Spot 故障恢復、vLLM TPU 功能覆蓋,以及在非 GKE Kubernetes 上的可攜性。

來源

  1. Run Ray on TPU, Part 2: Ray AI libraries
  2. Get started with Ray on TPU with GKE