让大模型为 TPU 写低层 Pallas 内核,难点先落在“代码能否编译并算对”,随后才是“能否跑得更快”。JAXBench 用 50 个 JAX 工作负载评测这件事:在 Gemini 3 Flash、相同迭代搜索流程和相同样本预算下,不给 TPU 专属资料时,候选内核的单样本正确率为 5.8%;把整理后的 TPU 架构说明、Pallas API 参考、代码示例和规则块放进每轮上下文后,这一指标升至 37.3%。这里的正确率按所有生成候选计算,48/50 则是迭代加上下文方法最终至少产出一个正确内核的工作负载数,两个数字不能混为一谈。
这组实验讨论的对象:Pallas 是 JAX 面向 TPU 的低层内核编程接口,生成结果需在 TPU v6e 上编译、以 BF16 输入做数值校验,并通过 profiler 计时。JAXBench 的价值在于把“代码写出来”与“是否真的快”放在同一套可复核流程里。结果显示,文档主要减少 API 和运行时层面的失败;当正确种子出现后,如何分配搜索预算才决定能把速度推进到什么程度。
[图1:JAXBench 概览:在 TPU v6e 上对 AI 生成的 Pallas 内核与手工调优基线进行评估的 TPU 原生基准。]

JAXBench 为什么单独做 TPU 内核评测
GPU 内核优化的经验不能直接套到 TPU。TPU v6e 使用宽 SIMD 向量寄存器和 256×256 的脉动矩阵乘法单元 MXU,执行模型与 GPU 的 SIMT 不同;JAX 程序经 XLA 编译,低层内核由 Pallas 降到 Mosaic 后端。内核作者需要处理 VMEM、SMEM、HBM 的内存层级、软件流水线、预取、块形状和网格遍历等约束。
现有基准还有另一个问题:如果问题规模太小,MXU 无法充分利用,测到的差距容易被启动和访存开销主导。JAXBench 从 MaxText 的公开模型中取出 17 个生产算子,又将 KernelBench Level 2 的 33 个融合算子转成 JAX;每个工作负载都调到让 XLA 基线至少达到 60% MXU 利用率。8 个优先算子另有 Tokamax 的手工 Pallas 实现,可用来观察自动方法与专家实现之间的距离。
[表1:LLM 内核生成基准对比。]

文档补上的,是 Pallas 的可执行约束
在无上下文的迭代流程中,Gemini 3 Flash 会拿到编译错误和运行反馈,却仍常在 Pallas API、pallas_call 参数和 BlockSpec 类型上出错。原始统计中,93.8% 的无上下文迭代候选在编译或首次执行阶段因 API 误用失败。很多与内存布局、可整除性和预取调度有关的约束,并不会完整出现在报错文本里,单靠继续试错很难恢复。
迭代加上下文方法在每轮提示前加入四类材料:硬件架构摘要、按基准挑选的 Pallas API 参考、精选代码示例和规则块。搜索算法没有改动,单样本正确率便从 5.8% 提升到 37.3%;最终 50 个工作负载中有 48 个得到过正确内核,几何平均加速为 XLA 的 1.28 倍。这里能够成立的判断是:在这套 Pallas 任务和 Gemini 3 Flash 条件下,目标资料显著改善了候选代码的正确性,不能据此推及所有 DSL 或所有模型。
[图4:JAXBench 上每个样本的结果分类(Gemini 3 Flash),汇总 50 个基准上的全部样本。]

写对之后,搜索预算才开始影响速度
迭代加上下文把更多样本用于排错,因此解决工作负载的数量最高,却不一定给每个正确内核留下足够的性能搜索空间。Autocomp 把预算拆为翻译与优化两阶段:先用束搜索形成正确种子,再把后续候选投入性能优化。在同样的 50 个工作负载和 Gemini 3 Flash 设置下,它最终达到 1.36 倍几何平均加速,高于迭代加上下文的 1.28 倍;后者在正确工作负载数上反而更多,为 48/50,对比 Autocomp 的 45/50。
这一区分很关键。48/50 描述的是“至少有一个候选算对”的覆盖面,1.36 倍描述的是把各工作负载最佳速度汇总后的整体结果。JAXBench 在汇总时把错误或缺失的工作负载按 1 倍计入,因此速度数字同时受正确率和已有正确内核的优化幅度影响。
[表2:Gemini 3 Flash 在完整 50 基准 JAXBench 套件上的跨方法对比。]

手工调优仍在分页注意力等长尾算子上占优
在有手工参考的 8 个优先算子子集上,Tokamax 的手工 Pallas 内核相对 XLA 达到 2.08 倍几何平均加速,Autocomp 为 1.60 倍,约为手工结果的 77%。自动方法在两个算子上超过该手工参考,并在另外四个算子上达到其 68% 至 91%;分页注意力和 ragged attention 仍未生成正确实现,而这两类算子尤其依赖手写调度。
这也给出 JAXBench 当前最清楚的结论:对 Pallas 这类训练资料稀少的接口,先提供可执行的硬件与 API 约束,能明显提高生成正确内核的机会;正确内核出现后,搜索结构负责把一部分正确性转成速度。全套评测只覆盖单颗 TPU v6e,多芯片分片和集合通信还未纳入,因此这些结果尚不能代表完整生产集群上的端到端收益。
[表3:8 个带有手工调优 Pallas 参考实现的优先内核上的加速。]

📄 原文标题
JAXBench: Benchmarking Autonomous TPU Kernel Optimization
🔗 原文链接
https://arxiv.org/abs/2607.20466
夜雨聆风