论文
JAXBench:自主TPU内核优化基准
JAXBench: Benchmarking Autonomous TPU Kernel Optimization
摘要
严格的基准通过确立共同的爬坡目标推动了自主GPU内核性能优化的进步,但TPU上没有对应物。我们提出JAXBench,一个面向谷歌云TPU上AI生成内核优化的TPU原生基准套件:含50个既相关又有优化余地的JAX工作负载。我们从公开MaxText库的架构(Llama-3.1、DeepSeek-V3、Mixtral、Mamba-2、AlphaFold2)抽取17个生产ML算子,翻译KernelBench的33个算子并验证正确性、设定 achieving高TPU v6e MXU利用率的规模;17个生产算子中8个附带公开Tokamax库的手工优化Pallas内核与块尺寸调优,以确立专家上限基线。我们在JAXBench上评估四种反馈驱动的候选Pallas内核生成方法。用Gemini 3 Flash跑全套件:在Pallas这类文档稀疏的DSL上,目标专属上下文比模型规模更重要——以精选TPU文档为条件把每样本正确率从5.8%提到37.3%、以1.28倍几何均值加速解决50个基准中的48个;正确性达成后搜索结构带来显著增益,Autocomp的束搜索管线相对XLA达1.36倍几何均值加速;在8个手工调优内核上Autocomp达1.60倍,收回2.08倍Tokamax上限的大部分,但在专用分页与不规则注意力算子上落后。高质量TPU内核优化仍是挑战性任务,我们发布JAXBench基准、评估环境与基线结果以支持开源贡献。
