桃子桃子快讯
返回首页
研究论文

JAXBench:面向 TPU 的 AI 内核优化基准发布

arXiv 新论文提出 JAXBench,首个针对 Google Cloud TPU 的 AI 生成内核优化基准,含 5…

2026.07.24 · 周五2 分钟阅读

arXiv 上一篇新论文提出了 JAXBench,这是首个专为 Google Cloud TPU 设计的 AI 生成内核优化基准套件,填补了 TPU 领域缺乏与 GPU KernelBench 对应基准的空白。论文同时开源了基准、评测框架与基线结果,以推动社区在 TPU 内核自动优化方向上的研究。

基准构成:50 个 JAX 工作负载

JAXBench 共包含 50 个 JAX 工作负载,覆盖范围兼具代表性与优化空间:

  • 生产级算子 17 个:从公开的 MaxText 模型库中提取,涵盖 Llama-3.1、DeepSeek-V3、Mixtral、Mamba-2 与 AlphaFold2 等主流架构中的真实 ML 算子。
  • 翻译算子 33 个:从 KernelBench 移植而来,经过正确性验证,并重新设置了能跑满 TPU v6e MXU 利用率的问题规模。
  • 专家基线:其中 8 个生产算子附带 Tokamax 库中手工优化的 Pallas 内核,并完成 block-size 调优,作为"专家上限"基线。

方法评估:四个反馈驱动方案

论文评估了四种反馈驱动的 Pallas 内核生成方法,整体使用 Gemini 3 Flash 作为生成模型。核心发现包括:

  • 在 Pallas 这种文档稀疏的 DSL 上,针对目标的上下文比模型规模更关键。
  • 加入经过整理的 TPU 文档后,单样本正确率从 5.8% 提升至 37.3%,并能在 50 个基准中解出 48 个,几何均值加速比为 1.28 倍。
  • 在正确性达到后,搜索结构带来显著增益:Autocomp 的束搜索(beam-search)流水线相对 XLA 达到 1.36 倍几何均值加速比。
  • 在 8 个手工调优内核上,Autocomp 相对 XLA 达到 1.60 倍几何均值加速比,已接近 Tokamax 设定的 2.08 倍上限,但在分页注意力(paged attention)与 ragged attention 等专用算子上仍有差距。

意义与开放问题

JAXBench 把 GPU 端 KernelBench 范式平移到 TPU,为自动内核优化研究提供了共享的评测目标。作者指出,高质量 TPU 内核优化仍是困难任务,期待社区通过开源贡献进一步完善基准与方法。

信源