現代 AI 基礎設施的建置挑戰已根本轉變。機器學習的現代前沿現已需要利用分散式系統,橫跨數千個加速器。隨著模型擴展至在數十萬個晶片的叢集上運行,驅動這些模型的軟體必須滿足效能、硬體可移植性和可靠性的新需求。

在 Google,我們的 Tensor Processing Units (TPUs) 是我們超級運算基礎設施的基石。這些客製化 ASIC 為 Gemini 和 Veo 等 Google 自有 AI 平台以及我們 Cloud 客戶的大規模工作負載提供訓練和服務。整個 AI 社群都應該能夠輕鬆存取 TPU 的全部功能,由於許多潛在使用者使用 PyTorch 建置模型,因此讓 PyTorch 能原生且高效地在 TPU 上運作的整合至關重要。

這就是 TorchTPU 的由來。作為一個工程團隊,我們的任務是建置一個以易用性、可移植性和卓越效能為首要考量的堆疊。我們希望讓開發者能夠以最少的程式碼變更來遷移現有的 PyTorch 工作負載,同時為他們提供提取我們硬體運算能力的 API 和工具。以下將深入探討驅動 TorchTPU 的工程原則、我們建置的技術架構,以及我們到 2026 年的發展藍圖。

要理解 TorchTPU,首先必須了解其目標硬體。

TPU 系統不僅僅是一個晶片;它是一個整合的網路。一個主機連接到多個晶片,每個晶片透過我們的 Inter-Chip Interconnect (ICI) 連接到主機和其他晶片。這個 ICI 將晶片連結成一個高效的 2D 或 3D 環形拓撲,能夠大規模擴展而無傳統網路瓶頸。在每個晶片內部,執行被劃分為 TensorCores 和 SparseCores。TensorCores 是專用於密集矩陣運算的單執行緒單元,而 SparseCores 則處理不規則記憶體存取模式,例如嵌入、gather/scatter 操作和集體通訊卸載。

這些功能意味著 TPU 是機器學習的強大工具;我們的目標是提供完全利用這些獨特功能所需的專門支援。這就是 PyTorch 的用武之地:PyTorch 工具鏈已經為其他裝置類型建立了統一、廣泛使用的介面。

我們易用性的核心原則很簡單:它應該感覺就像 PyTorch。開發者應該能夠採用現有的 PyTorch 腳本,將初始化更改為「tpu」,然後在不修改核心邏輯的任何一行程式碼的情況下運行其訓練迴圈。

要實現這一點,需要一種全新的方法來處理 PyTorch 與 TPU 編譯器和運行時堆疊的互動方式。

從概念到 TPU 上的原生 PyTorch 體驗,意味著需要重新思考執行堆疊。我們確立了「Eager First」的哲學。我們沒有要求開發者立即進行靜態圖編譯,而是透過 PyTorch 的「PrivateUse1」介面實現了 TorchTPU。沒有子類別,沒有包裝器;只有 TPU 上普通、熟悉的 PyTorch Tensor。透過在如此底層進行整合,我們能夠完全優先考慮開發者期望從 PyTorch 獲得的 Eager 執行體驗。

我們設計了三種不同的 Eager 模式來支援開發生命週期。

第一種 Eager 模式是 Debug Eager,它一次調度一個操作,並在每次執行後與 CPU 同步。它本質上很慢,但對於追蹤形狀不匹配、NaN 值和記憶體不足崩潰非常寶貴。

第二種是 Strict Eager,它保持單一操作調度,但異步執行,旨在鏡像預設的 PyTorch 體驗。這允許 CPU 和 TPU 同時執行,直到使用者腳本中的同步點。

然而,突破在於我們的 Fused Eager 模式。透過對操作串流進行自動反射,TorchTPU 在將其交給 TPU 之前,會將步驟即時融合為更大、計算密集型的區塊。透過最大化 TensorCore 利用率並最小化記憶體頻寬開銷,Fused Eager 在無需使用者進行任何設定的情況下,持續提供比 Strict Eager 高 50% 到 100% 以上的效能提升。

所有這三種模式都由一個共享的編譯快取支援,該快取可以在單一主機上運行,或配置為跨多主機設定持久化。這意味著,隨著 TorchTPU 學習您的工作負載,您將花更少的時間進行編譯,而花更多時間運行。

對於希望解鎖 TPU 最高效能的使用者,TorchTPU 與 torch.compile 介面原生整合,以實現全圖編譯。我們首先使用 Torch Dynamo 捕捉 FX 圖。然而,我們沒有透過 Torch Inductor 路由,而是使用 XLA 作為我們的主要後端編譯器。

這是一個經過深思熟慮的架構決策。XLA 已經過 TPU 拓撲的嚴格實戰測試。更重要的是,它原生了解如何優化密集計算與 ICI 上的集體通訊之間的關鍵重疊。我們的轉換層將 PyTorch 的運算子直接映射到 StableHLO,這是 XLA 的主要張量數學中間表示 (IR)。這在 PyTorch 和 XLA 的核心降低路徑之間建立了一個直接連接,使我們能夠生成高度優化的 TPU 二進位檔,同時重用我們 Eager 模式所建立的執行路徑。

對於編寫自訂運算子的開發者,我們確保可擴展性不會破壞效能。TorchTPU 原生支援使用 Pallas 和 JAX 編寫的自訂核心。透過使用 @torch_tpu.pallas.custom_jax_kernel 裝飾 JAX 函數,工程師可以編寫與我們的降低路徑直接互動、低階的硬體指令。目前也正在努力支援 Helion 核心。

為了在規模化時保留 Eager 和編譯模式的靈活性和易用性,我們大力關注 PyTorch 的分散式 API。如今,TorchTPU 開箱即支援 Distributed Data Parallel (DDP)、Fully Sharded Data Parallel v2 (FSDPv2) 和 PyTorch 的 DTensor。我們已經驗證,許多基於 PyTorch 分散式 API 建置的第三方函式庫在 TorchTPU 上無需更改即可運行。

PyTorch/XLA(TorchTPU 的前身)的一個主要限制是它只支援純 SPMD 程式碼。PyTorch 輸入的現實情況是,不同 rank 上運行的程式碼經常存在細微差異:例如,「rank 0」進程為了記錄或分析而執行額外工作是很常見的。這種輸入對 TPU 堆疊構成挑戰,而 TPU 堆疊針對 SPMD 優化進行了高度優化。XLA 在處理系統上運行的程式碼的全局視圖時效果最佳,但繞過它會增加開發者的開銷,他們必須小心地移除不純粹的行為。

TorchTPU 的架構旨在仔細支援發散式執行 (MPMD),並在必要時隔離通訊原語,以最低成本保持正確性。這種方法有助於確保在 TPU 上使用 PyTorch 的體驗對現有的 PyTorch 開發者盡可能自然,同時在可能的情況下,保留 XLA 透過分散式 TPU 部署的全局視圖來重疊通訊和計算的能力。

TPU 可以實現非常高的效能和效率,但最佳模型設計可能與其他硬體略有不同。例如,我們經常看到模型將 attention head 維度硬編碼為 64,而當前一代 TPU 在 128 或 256 的維度上實現了最高的矩陣乘法效率。修改模型以目標為 128 或 256 維度,可以更好地利用 TPU 晶片上龐大、密集且高效的張量核心。

可移植性並不能消除硬體現實,因此 TorchTPU 促進了分層工作流程:首先建立正確的執行,然後使用我們即將推出的深度指導方針來識別和重構次優架構,或注入自訂核心,以實現最佳硬體利用率。

我們今天已經為訓練和服務支援奠定了堅實的基礎,並且我們正在積極解決幾個開放性挑戰,以使 TorchTPU 成為 PyTorch 生態系統中無摩擦的後端。

我們編譯器團隊的一個主要重點是減少由動態序列長度和批次大小觸發的重新編譯。透過在 XLA 中實現先進的有界動態性,我們的目標是在不產生編譯開銷的情況下處理形狀變化。這對於某些工作負載(例如迭代式下一個 token 預測)來說可能是一個重要功能。

我們還正在為標準運算子建置一個全面的預編譯 TPU 核心函式庫,以大幅減少第一次迭代執行的延遲。

展望 2026 年的剩餘時間,我們正在努力:

TorchTPU 代表了我們致力於在 TPU 硬體上提供無縫、高效能 PyTorch 體驗的工程努力。我們正在打破障礙,消除您喜愛的框架與下一代 AI 所需的 TPU 超級運算硬體之間的摩擦。

如欲隨時了解最新的 TorchTPU 更新,請造訪 TPU 開發者中心。

在 TPU 上運行 Ray,第 2 部分:Ray AI 函式庫

在 TPU 上運行 Ray,第 1 部分:基礎