发表机构
Harvard University; Google; UC Berkeley; Google DeepMind(哈佛大学; 谷歌公司; 加州大学伯克利分校; 谷歌DeepMind)
机构由 AI 辅助整理,请以论文原文为准。AI 中文总结
该研究针对TPU缺乏内核优化基准测试的问题,提出JAXBench套件,包含多个工作负载和算子。通过评估四种反馈驱动方法,发现特定上下文对优化更关键,条件设定和搜索结构能提升性能,最后发布相关基准测试等支持开源。
AI 中文摘要
严格的基准测试通过建立共同目标推动了自主GPU内核性能优化的进展,但TPU却没有类似的测试。我们提出了JAXBench,这是一个用于在谷歌云TPU上进行人工智能生成的内核优化的原生TPU基准测试套件。它包含50个相关且有优化空间的JAX工作负载。我们从公共MaxText库的架构中提取了17个生产ML算子,从KernelBench中翻译了33个算子并设置新问题规模。我们评估了四种反馈驱动方法。结果表明特定上下文比模型规模更重要,条件设定提高了正确率,搜索结构在正确性达成后带来显著提升。我们发布JAXBench基准测试、评估工具和基线结果以支持开源贡献。
英文摘要
Rigorous benchmarks have driven progress in autonomous GPU kernel performance optimization by establishing a shared target to hillclimb on, but no equivalent exists for TPUs. We present JAXBench, a TPU-native benchmark suite for AI-generated kernel optimization on Google Cloud TPUs. JAXBench comprises 50 JAX workloads that are both relevant and provide headroom for optimization. We extract 17 production ML operators from architectures in the public MaxText library such as Llama-3.1, DeepSeek-V3, Mixtral, Mamba-2, and AlphaFold2, and translate 33 operators from KernelBench that are validated for correctness and set with new problem sizes that achieve high TPU v6e MXU utilization. Eight of the 17 production operators ship with hand-optimized Pallas kernels from the public Tokamax library and block-size tuned to establish an expert upper-bound baseline. We evaluate four feedback-driven methods on generating candidate Pallas kernels for JAXBench. Across the full suite with Gemini 3 Flash, we find that target-specific context matters more than model scale on a sparsely-documented DSL like Pallas. Conditioning on curated TPU documentation raises per-sample correctness from 5.8% to 37.3% and solves 48 of 50 benchmarks at a 1.28x geomean speedup. Search structure yields significant gains once correctness is achieved, with Autocomp's beam-search pipeline reaching a 1.36x geomean speedup over XLA. On the 8 hand-tuned kernels, Autocomp reaches 1.60x geomean over XLA, recovering most of the 2.08x Tokamax upper bound but trailing on the specialized paged and ragged attention operators. High-quality TPU kernel optimization remains a challenging task, and we release the JAXBench benchmark, evaluation harness, and baseline results to support open source contributions.