返回首頁

推論系統

Inferact 開源 TPU 推論巨核心,Kimi K3 低併發解碼逾每秒 700 token

實作以單次 Pallas 呼叫串起 92 層,透過跨層預取權重提高記憶體頻寬利用率。官方比較限於特定拓撲與小批次,推測解碼成績也取決於草稿接受率。

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

Inferact 於 9 月 23 日公開面向 TPU v7 的推論巨核心,以 Pallas 實作 Kimi K3 解碼。官方在 16 顆 TPU 上測得單一使用者每秒逾 700 個輸出 token;對照為 16 張 GB200、採 vLLM 配方的每秒 452 個。圖表使用推測解碼、接受長度為 6,成績應連同這項條件解讀。[官方技術文章](https://inferact.ai/blog/tpu-megakernels)

這套實作將 92 個 MoE 層合併到一次 Pallas 呼叫,在目前層進行專家運算時,預取下一層的注意力投影權重。目的在於讓運算、通訊與權重搬移重疊,減少解碼時高頻寬記憶體介面閒置。這對小批次服務特別有意義:每次只生成少量 token,仍須反覆搬入大量模型權重。[核心設計說明](https://inferact.ai/blog/tpu-megakernels)

Pallas 的控制方式是理解這項最佳化的關鍵。JAX 文件說明,TPU 主要依序執行指令,但可讓 DMA 搬移與矩陣運算在背景進行;核心使用片上 VMEM 等緩衝區,減少直接等待 HBM 的時間。工程師因而能安排資料存活時間與預取順序,但緩衝區配置超出容量時仍會編譯失敗,融合本身並不保證加速。[Pallas 文件](https://docs.jax.dev/en/latest/pallas/tpu/details.html)

這也提供編譯器研究的具體案例:跨層資料搬移的先後順序,可能比單一算子的局部最佳化更左右結果。依此設計推論,若換用其他模型,權重尺寸、注意力狀態與暫存需求都會改變,原有排程未必能直接沿用;重現者需要同時檢查數值正確性與記憶體排程,才能判斷收益來自哪個環節。

重現時還須分清實體晶片與框架裝置。Google 文件列明,每顆 Ironwood 會向 JAX 暴露兩個裝置,16 顆晶片的配置分布於四部主機。因此儲存庫寫的 32 個 TPU 裝置,與文章的 16 顆晶片相符;部署者應核對拓撲,不能只比較裝置數量。[Google 架構文件](https://docs.cloud.google.com/tpu/docs/tpu7x)

公開程式提供互動展示、HTTP 服務及正確性測試,但 README 明確指出,CPU 測試不涵蓋 TPU 編譯、VMEM 容量、DMA 排程與效能。實作也針對 2×2×4 拓撲,高併發下通訊與向量運算可能成為新瓶頸。對服務團隊而言,下一步應固定模型版本與草稿接受率,再量測長上下文、首 token 延遲和多請求負載;目前數字不足以推算完整服務成本。[程式與測試範圍](https://github.com/Inferact/tpu-megakernels)、[適用限制](https://inferact.ai/blog/tpu-megakernels)

來源

  1. 700 TPS on Kimi K3: A Case for TPU Megakernels
  2. Inferact/tpu-megakernels
  3. Writing TPU kernels with Pallas
  4. TPU7x (Ironwood)