JAXBench:对自主TPU内核优化进行基准测试

arXiv cs.AI 论文

摘要

JAXBench是一个新的基准测试套件,包含50个JAX工作负载,用于评估在谷歌云TPU上的AI生成内核优化,提供人工调优基线和智能体评估框架。论文发现,基于精选TPU文档进行条件化处理能显著提升正确性和加速比,其中Autocomp波束搜索在人工调优内核上相比XLA实现了高达1.6倍的几何平均加速。

arXiv:2607.20466v1 公告类型:新论文 摘要:严格的基准测试通过建立共同的山峰攀登目标,推动了自主GPU内核性能优化的进步,但TPU领域尚无此类基准。我们提出JAXBench,这是一个面向TPU的原生基准测试套件,用于在谷歌云TPU上进行AI生成的内核优化。JAXBench包含50个JAX工作负载,这些工作负载既相关又具有优化空间。我们从MaxText公共库(如Llama-3.1、DeepSeek-V3、Mixtral、Mamba-2和AlphaFold2)中的架构中提取了17个生产级ML算子,并将KernelBench中的33个算子转换过来,验证其正确性并设置新的问题规模,以实现TPU v6e MXU的高利用率。17个生产级算子中有8个附带了来自公共Tokamax库的手工优化Pallas内核,并调整了块大小以建立专家上限基线。我们评估了四种反馈驱动方法,用于生成JAXBench的候选Pallas内核。在使用Gemini 3 Flash的完整套件中,我们发现对于像Pallas这样文档稀疏的DSL,目标特定的上下文比模型规模更重要。基于精选TPU文档进行条件化处理,将每个样本的正确性从5.8%提升到37.3%,并在1.28倍几何平均加速下解决了50个基准中的48个。一旦正确性得以实现,搜索结构即可带来显著收益,Autocomp的波束搜索流水线相比XLA达到了1.36倍的几何平均加速。在8个人工调优内核上,Autocomp相比XLA达到了1.60倍的几何平均加速,恢复了对2.08倍Tokamax上限的大部分性能,但在专门的分页和稀疏注意力算子方面仍落后。高质量的TPU内核优化仍是一项具有挑战性的任务,我们发布了JAXBench基准测试、评估框架和基线结果,以支持开源贡献。
查看原文
查看缓存全文

缓存时间: 2026/07/24 05:00

# JAXBench:自主 TPU 内核优化基准测试

来源:https://arxiv.org/html/2607.20466

\paperurl\uselogo\correspondingauthor aryatschand@g\.harvard\.edu, sethusankaran@google\.com

Charles HongEqual contributions, work done while at GoogleUC BerkeleyGoogleJulian WalkerGoogle DeepMindNina CaiGoogleShangkun WangGoogleSuvinay SubramanianGoogleSundar DevGoogleVijay Janapa ReddiHarvard UniversityAmir YazdanbakhshGoogle DeepMindSethu SankaranGoogle

###### 摘要

代码:https://github.com/AI-Hypercomputer/accelerator-agents/tree/main/JAXBench

严格的基准测试通过建立共同的目标来推动 GPU 内核自主性能优化的发展,但 TPU 尚缺少类似的基准。我们提出 **JAXBench**,一个面向 Google Cloud TPU 上 AI 生成内核优化的 TPU 原生基准测试套件。**JAXBench** 包含 50 个 JAX 工作负载,这些工作负载既**相关**,又为优化提供了**提升空间**。我们从公共 MaxText 库(如 Llama-3.1、DeepSeek-V3、Mixtral、Mamba-2 和 AlphaFold2)的架构中提取了 17 个生产级 ML 算子,并从 KernelBench 翻译了 33 个算子,这些算子经过正确性验证,并设置了新的问题规模,以在 TPU v6e 上实现高 MXU 利用率。其中 8 个生产算子附带了来自公共 Tokamax 库的手工优化 Pallas 内核,并调整了块大小,以建立专家上限基线。我们评估了四种基于反馈的方法,用于为 **JAXBench** 生成候选 Pallas 内核。在完整套件上使用 Gemini 3 Flash,我们发现,对于像 Pallas 这样文档稀疏的 DSL,特定于目标的上下文比模型规模更重要。基于整理的 TPU 文档进行条件设置,将每个样本的正确率从 5.8% 提高到 37.3%,并在 50 个基准测试中解决了 48 个,几何平均加速比为 \(1.28\times\)。一旦实现正确性,搜索结构会带来显著的收益,Autocomp 的束搜索流水线相比 XLA 实现了 \(1.36\times\) 的几何平均加速比。在 8 个手工调优的内核上,Autocomp 相比 XLA 实现了 \(1.60\times\) 的几何平均加速比,基本达到 \(2.08\times\) 的 Tokamax 上限,但落后于专门的页式注意力和不规则注意力算子。高质量的 TPU 内核优化仍然是一项具有挑战性的任务,我们发布 **JAXBench** 基准测试、评估框架和基线结果,以支持开源贡献。

## 1 引言

高效的内核实现是充分发挥机器学习硬件加速器潜力的瓶颈。诸如 cuBLAS、CUTLASS 和 Triton\[cutlass,tillet2019triton\] 等库,以及针对 GPU\[dao2022flashattention\] 和 TPU\[jiang2026ragged\] 的专用内核,已经实现了逐代代的性能提升,但每种新的模型架构、量化方案和硬件版本都需要新的底层实现。要释放架构带来的性能优势,必须根据硬件特定功能和新兴工作负载特征,快速协同设计内核。

参考图注
图 1:JAXBench 概览:一个 TPU 原生基准测试,用于评估 AI 生成的 Pallas 内核与 TPU v6e 上手工调优基线的比较。我们构建了 50 个 JAX 参考工作负载(蓝色),包括来自 MaxText\[maxtext2022\] LLM 的 17 个优先级内核和改编自 KernelBench L2 的 33 个融合算子,所有算子的规模都足以使 TPU v6e MXU 饱和。我们还从 Tokamax 中提取了 8 个优先级内核的手工调优 Pallas 实现(绿色)作为专家上限。代理评估框架(紫色)能够评估任何 LLM 驱动的方法,以生成候选 Pallas 内核,这些内核经过编译、在 bf16 输入上进行正确性检查,并通过 `jax.profiler` 进行分析,反馈结果用于最终内核与两个基线的比较。

这一差距促使了使用语言模型自动生成和优化内核的快速增长的研究工作,范围从一次性完成到带有编译和分析反馈的迭代代理,再到程序空间上的进化搜索\[liao2025kernelevolve,tschand2025swizzleperf,lange2025towards,cao2026k\]。严格的基准测试一直是这一进展的核心。KernelBench\[ouyang2025kernelbench\] 通过 250 个 PyTorch 工作负载标准化了评估协议,并已扩展到 Triton、CuTe 和其他 DSL。TritonBench\[li2025tritonbench\] 专门针对 Triton 生成,并在 NVIDIA 和 AMD GPU 上进行硬件感知的性能测量。FlashInfer-Bench\[xing2026flashinfer\] 将评估基于真实的 LLM 服务轨迹,并使用专家编写的 FlashInfer 内核作为参考。这些基准测试共同推动了社区在正确性和相对于 GPU 基线的加速比方面的快速进步,并使从最佳 N 采样到强化学习编码代理等方法的有意义比较成为可能。然而,**TPU 内核优化尚无类似的基准测试**。

Google 的张量处理单元\[jouppi2017datacenter\] 与 GPU 有根本区别。TPU 是顺序机器,具有宽 SIMD 向量寄存器(v6e 上 32 位值为 \(8 \times 128\))和专用的 \(256 \times 256\) 脉动矩阵乘单元(MXU),而不是 GPU 的大规模并行 SIMT 执行模型。编程 TPU 还需要不同的软件栈。JAX\[bradbury2021jax\] 程序通过 XLA 编译,底层内核编写使用 Pallas,它降级到 Mosaic 后端而不是 Triton。Pallas 内核必须考虑 TPU 特定的问题,包括 VMEM/SMEM/HBM 内存层次结构、带有预取调度的软件流水线、块形状约束以及 Mosaic 强制执行的字典序网格遍历顺序\[jax\_pallas\_tpu\]。Pallas 在 LLM 训练数据中出现的频率也比 CUDA 甚至 Triton 低几个数量级,因此能够流畅编写 GPU 内核的模型经常会出现 Pallas API 幻觉、发出无法通过类型检查的内存空间注释,或者违反脉动平铺约束,这些问题是任何通用编译反馈都无法解决的。因此,GPU 基准测试无法直接应用,因为工作负载、问题规模、编程抽象和优化策略都特定于 GPU。

最近的并行工作解决了多平台评估的问题。MultiKernelBench\[wen2025multikernelbench\] 将 KernelBench 扩展到 CUDA、华为 AscendC 和 Google Pallas,发现最佳模型在 Pallas 任务上仅能达到 8.4–10.5% 的 Pass@1。然而,MultiKernelBench 不针对 JAX,在 TPU v2-8 上评估,将 PyTorch 工作负载翻译成问题规模太小而无法使 MXU 饱和,并且不包括生产级 LLM 算子或优化的参考内核。在小规模下,工作负载主要由内存流量和启动开销主导,而不是计算,因此任何测量的加速比都反映了簿记工作,而不是真正的算法或调度改进。一个有用的 TPU 基准测试必须将工作负载推入计算受限区域,其中 MXU 是瓶颈,而平铺、流水线和布局选择实际上会影响吞吐量。表 1 (https://arxiv.org/html/2607.20466#S1.T1) 总结了现有基准测试在这些轴上的定位。仍然需要一个 TPU 原生基准测试,它 (i) 使用当代 TPU 硬件,(ii) 包含代表现代 LLM 训练和推理的工作负载,(iii) 通过 XLA 编译和手工优化的 Pallas 内核提供强大的基线,以及 (iv) 将问题规模设为计算受限,以便优化空间是有意义的。

我们引入 **JAXBench**,一个专门为 TPU 内核优化设计的包含 50 个 JAX 工作负载的基准测试套件。我们的贡献如下:

1.  **包含 50 个相关工作负载的 TPU 原生基准测试套件。** **JAXBench** 是一个基准测试,包含来自真实 LLM 架构(Llama-3.1、DeepSeek-V3、Mixtral、Mamba-2、AlphaFold2)的 17 个生产算子,这些算子从 MaxText 中提取;以及 33 个来自 KernelBench 的融合算子序列,从 PyTorch 翻译为 JAX,并调整问题规模以实现高 TPU MXU 利用率。对于 8 个优先级内核,我们还提供了来自上游 Tokamax 库的专家优化 Pallas TPU 内核,并调整了块大小。
2.  **严格且可重复的评估框架。** 通过 `jax.profiler` Perfetto 跟踪进行设备端性能分析,消除了主机调度开销,提供了可重复的每迭代内核计时,可用于轻松评估任何在 **JAXBench** 上的代理框架。
3.  **代理评估和失败模式。** 我们在 TPU v6e(Trillium)上部署 **JAXBench**,并评估最佳 N 采样、迭代反馈代理循环、TPU 上下文条件代理循环,以及增强有 TPU 架构文档的 Autocomp\[hong2025autocomp\]。我们还提供了关于失败模式的见解,以及进一步贡献于以 TPU 为中心的内核优化的机会。

表 1:LLM 内核生成基准测试的比较。**JAXBench** 是唯一针对当代 TPU 的套件,具有生产规模的工作负载、硬件饱和的问题规模,以及针对子集的专家编写的 Pallas 参考内核。

| 基准测试 | 目标 | # 任务 | 工作负载来源 | 专家参考 | 饱和规模 |
|---|---|---|---|---|---|
| KernelBench\[ouyang2025kernelbench\] | CUDA / GPU | 250 | PyTorch 模块 | — | 是 |
| TritonBench\[li2025tritonbench\] | Triton / GPU | 184 | 真实仓库 | Triton | 部分 |
| FlashInfer-Bench\[xing2026flashinfer\] | CUDA / GPU | – | LLM 服务轨迹 | FlashInfer | 是 |
| MultiKernelBench\[wen2025multikernelbench\] | Pallas / TPU v2-8 | 285 | PyTorch 模块 | — | 否 |
| **JAXBench**(我们的) | Pallas / TPU v6e | 50 | 生产 LLM + 标准 JAX | Pallas | 是 |

## 2 方法论

### 2.1 设计原则

**JAXBench** 围绕三个原则构建,这些原则共同使基准测试对于 TPU 内核优化既具有实际相关性,又具有技术严谨性。

#### 相关的工作负载。
基准测试必须评估一组多样化的算子,这些算子代表真实世界的 TPU 用户。我们既包括已经存在专家优化的 Pallas 内核的经过充分研究的算子(例如,flash attention、grouped-query attention),也包括尚未开发出此类 Pallas 内核的新兴算子(例如,Mamba-2 状态空间对偶性、AlphaFold2 三角形乘法)。这种组合测试了自动化优化方法在不同人类先验努力水平下的表现。

#### 优化空间。
工作负载和问题规模的选择使得 XLA 编译的基线已经实现了高 TPU 利用率,让内核优化在真正的算法和调度改进上进行竞争,而不是填补人为的差距。这种选择反映了 TPU 的实际使用方式。运行生产工作负载的客户会调整问题维度以最大化利用硬件,而工作负载级别的 FLOPs 利用率与底层操作的 FLOPs 利用率紧密相关。一个在规模不足、内存受限或启动开销受限的规模上评估内核的基准测试,将无法反映 TPU 用户(以及 TPU 内核优化器)实际操作的场景。因此,我们将每个工作负载的规模设定为推动相关操作进入计算受限区域,其中 MXU 是瓶颈资源,优化空间是有意义的。如图 2(a) (https://arxiv.org/html/2607.20466#S2.F2.sf1) 所示,重矩阵乘的融合算子使用 bf16 中的维度如 \( (4096,8192) \times (8192,8192) \),达到 60–95% 的 MXU 利用率,优先级内核使用其源架构的生产规模维度。例如,Llama-3.1-70B GEMM 工作负载在 \(8192 \times 8192 \times 28672\) 上运行,达到约 79% 的 MXU 利用率。问题规模本身就是一个优化轴,因为不同的形状暴露了不同的机会,而足够小的问题根本没有优化空间。

#### 可重复的基线。
每个工作负载都使用惯用的 JAX 实现,通过 XLA 的 `jax.jit` 编译。这是 TPU 最强大的可重复基线,因为 XLA 已经执行了激进的整个程序优化,包括算子融合、缓冲区分配和布局优化。对于优先级内核,在可用的情况下,我们还提供 Pallas 优化的变体,代表手工调优 TPU 内核性能的当前最先进水平,建立自动化方法应努力达到的上限。

### 2.2 工作负载构建

#### 优先级内核(17 个工作负载)。
我们提取了 17 个来自生产 LLM 架构的计算关键算子,这些算子按照 MaxText\[maxtext2022\] 中的实现。该集合涵盖了注意力变体(flash\[dao2022flashattention\]、GQA\[grattafiori2024llama\]、MLA\[liu2024deepseek\]、稀疏/splash\[jiang2024mixtral\]、flex、页式\[kwon2023efficient\] 和不规则页式)、密集和稀疏线性代数(GEMM、SwiGLU MLP\[shazeer2020glu\]、稀疏 MoE\[shazeer2017outrageously\]、Megablox GMM、不规则点积)、归一化和损失(RMSNorm\[zhang2019root\]、交叉熵),以及新兴架构(RetNet 保留\[sun2023retentive\]、Mamba-2 SSD\[gu2023mamba\]、AlphaFold2 三角形乘法\[jumper2021highly\])。每个工作负载都使用其源模型的生产规模维度。例如,GQA 注意力工作负载使用 Llama-3.1-405B 维度,具有 128 个查询头和 8 个键值头,序列长度为 4096。在可用的情况下,我们通过从上游 JAX Pallas 算子库(Tokamax)\[tokamax2024\] 导入等效的 Pallas 内核来获取优化实现,然后通过 TPU v6e 上的穷举网格搜索调整其块大小参数。库中提供了八个这样的高质量内核。调优过程总共评估了 203 种配置,相比 Pallas 默认参数,加速比高达 \(2.79\times\)(Megablox GMM)。

#### KernelBench 融合算子(33 个工作负载)。
我们从 KernelBench 第 2 级\[ouyang2025kernelbench\] 翻译了一组精选的 33 个工作负载,从 PyTorch 转换为等效的 JAX。我们专注于第 2 级,因为第 1 级由孤立的单个算子组成,没有融合结构,而第 2 级将矩阵乘法或卷积与元素操作(激活、归一化、池化)结合起来,直接代表了与 TPU 优化最相关的内核融合机会。为了选择一个非冗余的子集,我们按融合签名对所有第 2 级任务进行分组,即它们实例化的(matmul 或卷积)×(激活、归一化、池化)算子类,每个类保留一个代表性任务。此过程消除了结构相同的任务,得到了我们包含的 33 个工作负载。翻译是通过使用每个参考 PyTorch 模块提示 Gemini,并请求一个惯用的 JAX 等效实现来执行的。为了确认数值正确性,我们在 TPU 上使用匹配的 bf16 输入执行原始 PyTorch 参考和生成的 JAX 实现,并要求它们的输出在 `jnp.allclose` 下一致,容差为 \( \mathrm{atol} = \mathrm{rtol} = 10^{-2} \),这是标准的 bf16 容差。任何未通过此检查的工作负载都会在包含之前重新生成或手动修复。许多原始 KernelBench 问题规模太小,无法使 TPU MXU 饱和,因为 \(256 \times 256\) 的脉动阵列需要大的矩阵维度才能实现高利用率。我们不应用统一的缩放因子,而是独立调整每个工作负载的自由维度,扫描候选形状,并冻结 XLA 基线在 TPU v6e 上至少达到 60% MXU 利用率的最小配置。这个按算子进行的过程产生了一组异构的形状(例如,矩阵乘融合算子在 bf16 中批量大小为 4096,特征维度为 \(8192 \times 8192\)),所有这些在 TPU 上都是计算受限的,而不是启动开销受限的。

#### 工作负载接口。
每个工作负载模块都公开一个标准化的接口,由一个 `CONFIG` 字典组成,其中包含超参数...

相似文章

TUA-Bench: 通用终端使用代理的基准测试

Hugging Face Daily Papers

TUA-Bench是一个综合性基准测试,用于评估通用终端使用代理在各种数字活动和专业工作流中的表现,揭示了当前前沿代理之间的显著性能差距。