论文

JAXBench:自主TPU内核优化基准

JAXBench: Benchmarking Autonomous TPU Kernel Optimization

智能体系统Agent任务评测

摘要

严格的基准通过确立共同的爬坡目标推动了自主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基准、评估环境与基线结果以支持开源贡献。

JAXBench:自主TPU内核优化基准:论文配图
图 1:JAXBench 概述:TPU 原生基准测试,用于根据 TPU v6e 上手动调整的基线评估 AI 生成的 Pallas 内核。我们构建了 50 个 JAX 参考工作负载(蓝色),其中包含来自 MaxText [4] LLM 的 17 个优先级内核和改编自 KernelBench L2 的 33 个融合运算符,所有大小都足以使 TPU v6e MXU 饱和。我们还从 Tokamax 中提取手工调整的 Pallas 实现(绿色)的 8 个优先级内核作为专家上限。代理评估工具(紫色)可以评估任何 LLM 驱动的方法,以生成候选 Pallas 内核,这些内核经过编译、对 bf16 输入进行正确性检查,并通过 jax.profiler 进行分析,反馈反馈并将最终内核与两个基线进行比较。