JAXBench:自主TPU核心最佳化的基準測試
JAXBench是一個專為Google Cloud TPU設計的基準測試套件,包含50個JAX工作負載,用於評估AI生成的核心最佳化方法。研究顯示,針對TPU的文件條件化能大幅提升正確性,而搜尋結構在正確性基礎上進一步加速,但高精度TPU核心最佳化仍有挑戰。
JAXBench是一個全新的TPU原生基準測試套件,旨在推動AI生成核心在Google Cloud TPU上的效能最佳化。該套件包含50個JAX工作負載,這些負載既具有實際相關性,又為最佳化提供了充足空間。其中,17個生產級機器學習運算元來自MaxText庫中的架構,如Llama-3.1、DeepSeek-V3、Mixtral、Mamba-2和AlphaFold2;另外33個運算元從KernelBench轉化而來,經過正確性驗證,並設定了新的問題規模以實現高TPU v6e MXU利用率。值得注意的是,17個生產運算元中有8個附帶了來自Tokamax庫的手工最佳化Pallas核心,並進行了塊大小調整,從而建立了專家上限基線。
研究團隊評估了四種基於反饋的方法用於生成候選Pallas核心。使用Gemini 3 Flash模型進行全面測試後,發現目標特定上下文比模型規模更為重要,尤其是在像Pallas這樣文件稀疏的DSL中。透過基於精選TPU文件的條件化,每個樣本的正確性從5.8%提升至37.3%,並在50個基準測試中解決了48個,幾何平均加速比達到1.28倍。一旦實現了正確性,搜尋結構便帶來了顯著增益:Autocomp的波束搜尋流水線實現了相對於XLA的1.36倍幾何平均加速比。在8個手工調優的核心上,Autocomp達到1.60倍幾何平均加速比,幾乎追上了Tokamax 2.08倍的上限,但在專門的分頁和稀疏注意力運算元上仍有差距。
雖然高精度的TPU核心最佳化仍是一項艱鉅任務,但研究團隊已經發布了JAXBench基準測試、評估框架和基線結果,以支援開源社群的進一步貢獻。這些工具為未來的研究工作提供了一個共享的最佳化舞臺,有望加速TPU核心自動最佳化的進展。研究指出,對於稀疏文件的DSL,提供有針對性的上下文資訊比單純擴大模型規模更為關鍵,這一發現對於其他類似領域的自動程式碼生成也具有指導意義。