返回首頁

AI 推論與編譯器

PyTorch Helion 新增 TPU 後端,以同一套高階核心程式跨接 Triton 與 Pallas

PyTorch 與 Google 為 Helion 建立 TPU 後端,讓 PyTorch 風格的核心程式可編譯成 Pallas,減少另寫 TPU 低階核心的負擔。團隊在 TPU v7 測得 Flash Attention 最高 838 TFLOPs,但功能仍處於積極開發階段。

Pytorch Deepdream (https://github.com/gordicaleksa/pytorch-deepdream) by gordicaleksa (https://github.com/gordicaleksa/pytorch-deepdream/commits?author=gordicaleksa) · MIT · Image source
zh-Hant

PyTorch 團隊 7 月 23 日宣布,機器學習核心 DSL Helion 已能以 Google Pallas 為後端產生 TPU 程式。Helion 原本主要把 Python 內嵌、接近 PyTorch 的張量與分塊程式編譯至 Triton;新增路徑後,工程師可望維護同一份高階核心,再依硬體選擇 GPU 或 TPU 後端,而不必為 TPU 直接重寫大量 Pallas 程式。

這項工作的重點不是單純換掉程式碼產生器。TPU 採少量大型工作單元及明確的 HBM、VMEM 記憶體層級,效能仰賴軟體安排非同步搬移,與 GPU 的 SIMT 模型不同。Helion 因此會把外層分塊轉為 `pallas_call` 管線,並在內層自動比較 `emit_pipeline` 與 `unroll`:前者逐塊把 K、V 從 HBM 搬入 VMEM,能處理較長序列;後者預先把資料留在 VMEM,可消除運算氣泡,但會隨序列長度增加記憶體需求。

官方在 B=8、H=32、D=256 的 Flash Attention 測試中,8K 序列使用 unroll 得到 892 TFLOPs,而 32K 序列因 VMEM 不足無法採用該策略;自動調校器可依形狀改選管線。另一項 TPU v7 測量則達 838 TFLOPs,約為單一 tensor core 的 79% MFU。較廣泛的核心集合中,團隊報告 Helion 相對 TorchTPU eager 幾何平均快 1.55 倍、相對 `torch.compile` 快 1.12 倍。

這使跨加速器核心的可攜性更接近實務需求,但數字來自專案團隊選定的核心、形狀與硬體,不能直接推論到完整模型訓練或推論。接下來應觀察 TPU 後端支援範圍、編譯與自動調校成本,以及同一份 Helion 核心在不同 TPU、GPU 世代上能否穩定取得接近手寫核心的效能。

來源

  1. Helion on TPU: Towards Hardware Heterogeneous Kernel Authoring
  2. pytorch/helion