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,提供有针对性的上下文信息比单纯扩大模型规模更为关键,这一发现对于其他类似领域的自动代码生成也具有指导意义。