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
Evaluation benchmarks have driven progress in automated kernel optimization, yet existing suites target GPUs exclusively. We present JAXBench, a TPU-native benchmark for AI-generated kernel optimization on Google Cloud TPUs. JAXBench comprises 50 JAX workloads, including 17 production LLM operators extracted from architectures such as Llama-3.1, DeepSeek-V3, Mixtral, Mamba-2, and AlphaFold2, and 33 fused operator sequences adapted from KernelBench. Eight of the 17 production operators ship with hand-optimized Pallas TPU kernels whose block sizes we tune via grid search, establishing reference baselines. We evaluate one-shot generation, iterative coding agents, the same iterative loop with TPU documentation injected, and a TPU-enabled Autocomp configuration augmented with the same 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. All baselines still trail hand-tuned performance, showing that TPU kernel generation remains open. We release the benchmark, harness, and baselines for reproducible research.
Chat is not available.
Successful Page Load