JAXBench: Benchmarking Autonomous TPU Kernel Optimization
Arya Tschand ⋅ Charles Hong ⋅ Julian Walker ⋅ Shengnan Cai ⋅ Shangkun Wang ⋅ Suvinay Subramanian ⋅ Sundar Dev ⋅ Vijay Janapa Reddi ⋅ Amir Yazdanbakhsh ⋅ Sethuraman Sankaran
Abstract
Rigorous benchmarks have driven progress in autonomous GPU kernel performance optimization by setting a shared target to hillclimb on, but no equivalent exists for TPUs. We present JAXBench, a TPU-native benchmark 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 LLM operators from architectures in the public MaxTech library such as Llama-3.1, DeepSeek-V3, Mixtral, Mamba-2, and AlphaFold2, and translate 33 fused operator sequences from KernelBench. Eight of the 17 production operators ship with hand-optimized Pallas TPU kernels from the public Tokamax library and block-size tuned for the target TPU. We evaluate one-shot generation, iterative coding agents, an iterative loop with TPU documentation injected, and a TPU-enabled Autocomp configuration augmented with the TPU-specific documentation. With Gemini 3 Flash, best-of-$N$ solves 13/50 at $1.01\times$ geomean and iterative refinement reaches 32/50 at $1.18\times$. Injecting TPU documentation lifts iterative refinement to 48/50 at $1.28\times$ and per-sample correctness from $5.8\%$ to $37.3\%$. Autocomp solves 45/50 and converts those correct kernels into $1.36\times$ geomean with $76\%$ beating XLA. High-quality TPU kernel optimization remains a challenging task, and we release the JAXBench benchmark, evaluation harness, curated documentation, and baseline results to support open source contributions.
Chat is not available.
Successful Page Load