AI News HubLIVE
站內改寫1 分鐘閱讀

JAXBench:自主TPU內核優化的基準測試

JAXBench是一個專為Google Cloud TPU設計的基準測試套件,包含50個JAX工作負載,用於評估AI生成的內核優化方法。研究顯示,針對TPU的文檔條件化能大幅提升正確性,而搜索結構在正確性基礎上進一步加速,但高精度TPU內核優化仍有挑戰。

來源arXiv AI作者: Arya Tschand, Charles Hong, Julian Walker, Nina Cai, Shangkun Wang, Suvinay Subramanian, Sundar Dev, Vijay Janapa Reddi, Amir Yazdanbakhsh, Sethu Sankaran

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,提供有針對性的上下文信息比單純擴大模型規模更為關鍵,這一發現對於其他類似領域的自動代碼生成也具有指導意義。