JAXBench: Benchmarking Autonomous TPU Kernel Optimization
Quick Answer
JAXBench introduces a TPU-native benchmark suite for optimizing AI-generated kernels on Google Cloud TPUs, featuring 50 JAX workloads.
Quick Take
It demonstrates a 1.28x speedup across benchmarks and a 1.60x speedup on hand-tuned kernels, highlighting the importance of target-specific context in kernel optimization.
Key Points
- JAXBench includes 50 relevant JAX workloads for TPU kernel optimization.
- Achieved a 1.28x speedup across benchmarks with improved correctness.
- Autocomp's beam-search pipeline yielded a 1.36x speedup over XLA.
- Eight hand-tuned kernels reached a 1.60x speedup compared to XLA.
- The benchmark suite supports open-source contributions for TPU optimization.
DeepSignal Analysis
What happened
JAXBench is a newly introduced benchmark suite designed for optimizing AI-generated kernels specifically on Google Cloud TPUs. It includes 50 JAX workloads and demonstrates significant performance improvements, achieving a 1.28x speedup across benchmarks and a 1.60x speedup on hand-tuned kernels.
Key evidence
- JAXBench comprises 50 JAX workloads relevant for optimization, extracted from production ML operators in the public MaxText library.
- The benchmark suite shows a 1.28x geometric mean speedup across all benchmarks and a 1.60x speedup on eight hand-tuned kernels compared to XLA.
- Conditioning on curated TPU documentation improved per-sample correctness from 5.8% to 37.3%, solving 48 out of 50 benchmarks.
Why it matters
The introduction of JAXBench fills a gap in TPU performance optimization, which previously lacked rigorous benchmarking like that available for GPUs. This suite not only aids in optimizing kernel performance but also emphasizes the importance of context-specific optimizations, potentially leading to more efficient AI model training and inference on TPUs.
What to watch
Paper Resources
📖 Reader Mode
~2 min readAbstract: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.
| Subjects: | Artificial Intelligence (cs.AI) |
| Cite as: | arXiv:2607.20466 [cs.AI] |
| (or arXiv:2607.20466v1 [cs.AI] for this version) | |
| https://doi.org/10.48550/arXiv.2607.20466 arXiv-issued DOI via DataCite (pending registration) |
Submission history
From: Arya Tschand [view email]
[v1]
Tue, 19 May 2026 05:38:39 UTC (614 KB)
— Originally published at arxiv.org
Want this in your inbox every morning?
Daily brief at your local 8am — bilingual EN/中文, free.
More from arXiv cs.AI
See more →HOBA: Hierarchical On-Policy Bidding Agents for Adaptive Online Advertising
HOBA (Hierarchical On-policy Bidding Agents) is a novel hierarchical reinforcement learning framework that enhances online advertising bidding systems by improving adaptability and reducing hyperparameter tuning costs. It utilizes a for hyperparameter inference, a SARSA agent for expert model selection, and a dynamic expert pool for bid execution, achieving a +3.6% increase in target cost during large-scale deployment and outperforming state-of-the-art baselines on AuctionNet.