TPU 内核优化终于有基准:喂文档比换大模型更管用,正确率 5.8% 到 37.3%

JAXBench: Benchmarking Autonomous TPU Kernel Optimization

Arya Tschand, Charles Hong, Julian Walker, Nina Cai, Shangkun Wang, Suvinay Subramanian, Sundar Dev, Vijay Janapa Reddi, Amir Yazdanbakhsh, Sethu Sankaran

cs.AI

2026-05-19

给 TPU 上的 Pallas 内核自动优化补了 50 项基准。对训练数据稀缺的 DSL,喂针对性文档比换更大模型更管用,单样本正确率从 5.8% 拉到 37.3%。

这篇在解决什么

GPU 上自动优化内核这件事有 KernelBench 这类基准当共同目标,TPU 上一直没有。TPU 是带矩阵乘单元(MXU)的序列机,要用的不是 CUDA 而是 Pallas 这个 DSL,而 Pallas 在训练语料里出现得少、文档也稀疏。没基准就没法衡量 LLM 生成的 Pallas 内核离人类专家多远,也没共同的山头让社区爬。Google 这篇就是来补这个空白,顺带问一个更具体的问题:对这种冷门 DSL,瓶颈到底是模型的推理能力,还是它根本没见过相关文档?

方法

JAXBench 50 项任务,全在 Google Cloud TPU v6e 单片上。构成:17 个生产算子,从 MaxText 里的真实架构抠出来(Llama-3.1、DeepSeek-V3、Mixtral、Mamba-2、AlphaFold2),覆盖各种 attention、GEMM、RMSNorm、MoE、RetNet 之类;33 个从 KernelBench Level 2 翻译过来,问题规模调到 MXU 利用率不低于 60%;8 个优先算子带 Tokamax 库里人工调过的 Pallas 内核当专家上界。

评测四档:能编译、正确性(jnp.allclose,bf16)、相对 XLA 的加速比、fast1@NN(N 次采样内压过 XLA 的比例)。加速比按每项最佳、floor 到 1 倍再取几何均值,错的算 1 倍。

测了四种反馈驱动方法(都用 Gemini 3 Flash,144 样本预算):Best-of-NN、迭代精修(带编译错误、正确性、profiler 反馈)、迭代精修加文档、Autocomp(两阶段 beam search,翻译 4 轮加优化 4 轮)。文档是一块针对性的 TPU 文档:硬件架构摘要、Pallas API 参考、代码示例、规则块。为什么这么设计?因为不喂文档时,99.7% 的 Best-of-NN 样本和 93.8% 的迭代样本在编译或首次执行就挂了,主要挂在 Pallas API 用错上。

结果

全 50 项,Gemini 3 Flash:

方法几何均值加速比正确数
Best-of-NN1.01 倍13/50
迭代精修1.18 倍32/50
迭代加文档1.28 倍48/50
Autocomp1.36 倍45/50

加文档把单样本正确率从 5.8% 干到 37.3%,48/50 能解。最干净的对照是:把模型从 Flash 升到 Pro(5 项子集),迭代几何均值从 1.18 倍到 2.43 倍;而加文档在 Flash 上把正确率拉了 31 个百分点。作者原话是换更大模型有帮助,但没加文档帮助大。结论是对冷门 DSL,瓶颈是信息不是推理。

在 8 个有人工内核的算子上,Autocomp 几何均值 1.60 倍,达到 Tokamax 上界 2.08 倍的约 77%。而且有反超:稀疏 attention Autocomp 2.81 倍对手调 0.86 倍,Megablox GMM 2.21 倍对 1.62 倍。栽跟头的是 paged 和 ragged paged attention,Autocomp 直接没产出正确内核(ragged 手调能到 6.91 倍),这正是手写调度最吃功力的地方。

为什么重要

对做编译器和内核自动优化的人,这篇的实操信号很硬:对训练数据稀缺的 DSL,先往 prompt 里塞针对性文档,比换更大的模型性价比高得多,这一条大概率也适用于其他冷门硬件 DSL。JAXBench 本身是个能用的山头,开放了基准、评测脚本和基线。更细的一条:解出更多题不等于更快,迭代加文档解得最多(48/50)但几何均值(1.28 倍)低于 Autocomp(1.36 倍),因为预算花在调试还是花在优化,效果不同。

局限与存疑

作者承认:全在单片 TPU v6e 上,多片 sharding 和集合通信没覆盖;人工内核只覆盖 17 个优先算子里的 8 个,上界比较是部分的;Pro 只在 5 项子集上跑(成本);paged 和 ragged attention 上 agent 完全解不了,作者自己说这仍是难题。另外几点:

术语

原文与代码

社区讨论

相关论文

全部论文解读