技术进展

JAXBench:TPU内核自动优化新基准

Heooo 07月24日12时31分 3 阅读

「JAXBench是专为Google Cloud TPU设计的AI内核优化基准测试套件,包含50个JAX工作负载,其中17个来自Llama-3.1、DeepSeek-V3等生产级架构,33个从KernelBench移植。研究显示,在Pallas这种文档稀疏的DSL上,针对特定目标的上下文比模型规模更重要,使用Gemini 3 Flash时,条件化TPU文档可将正确率从5.8%提升至37.3%,并解决48个基准,几何平均加速1.28倍。Autocomp的波束搜索管道达到1.36倍几何平均加速,在手调内核上恢复大部分Tokamax上限。该基准和评估工具已开源,旨在推动TPU内核优化社区发展。」

在AI硬件加速领域,GPU内核的自动优化已通过KernelBench等基准测试取得显著进展,但TPU(张量处理单元)的类似工具却长期缺失。近日,一篇题为《JAXBench: Benchmarking Autonomous TPU Kernel Optimization》的论文填补了这一空白,发布了专为Google Cloud TPU设计的AI生成内核优化基准测试套件JAXBench。该套件包含50个JAX工作负载,覆盖生产级ML算子与经典计算内核,为研究社区提供了统一的性能优化目标。

JAXBench的构建基于两大来源。首先,研究团队从MaxText公共库中提取了17个生产级ML算子,这些算子来自Llama-3.1、DeepSeek-V3、Mixtral、Mamba-2和AlphaFold2等知名架构。例如,Llama-3.1的注意力机制算子、DeepSeek-V3的混合专家层算子,以及AlphaFold2的结构预测算子,均被纳入测试集。这些算子不仅具有实际应用价值,还为优化提供了充足空间。其次,团队从KernelBench中移植了33个算子,经过正确性验证并调整了问题规模,以确保在TPU v6e上实现高MXU(矩阵乘法单元)利用率。这种混合设计使JAXBench既能评估TPU在真实AI负载下的表现,又能与GPU基准保持一定可比性。

为建立专家级性能上限,研究团队为17个生产算子中的8个提供了手调Pallas内核。这些内核来自公共Tokamax库,经过块大小调优,实现了2.08倍的几何平均加速比(相对于XLA编译器基线)。Pallas是Google为TPU开发的低级DSL(领域特定语言),允许开发者编写自定义内核以绕过XLA的自动优化限制。然而,Pallas的文档相对稀疏,这给AI模型生成正确内核带来了挑战。

论文评估了四种反馈驱动方法在JAXBench上生成候选Pallas内核的能力。使用Gemini 3 Flash模型时,研究揭示了一个关键发现:在Pallas这种文档稀疏的DSL上,针对特定目标的上下文信息比模型规模更重要。当模型仅依赖通用知识时,每个样本的正确率仅为5.8%。但通过条件化TPU文档(即提供Pallas语法、TPU硬件特性等针对性资料),正确率跃升至37.3%,并成功解决了50个基准中的48个,几何平均加速比达到1.28倍。这表明,对于专业领域代码生成,高质量文档的引导作用远超模型参数量的增加。

一旦实现正确性,搜索策略的优化效果显著提升。Autocomp的波束搜索管道在JAXBench上达到1.36倍的几何平均加速比,优于XLA基线。在手调的8个内核上,Autocomp达到1.60倍几何平均加速比,恢复了Tokamax专家上限的大部分性能,但在分页注意力(paged attention)和稀疏注意力(ragged attention)等特殊算子上的表现仍落后于人工调优。这反映出TPU内核优化仍存在挑战,尤其在处理不规则内存访问模式时,自动搜索难以完全匹配专家经验。

JAXBench的发布为TPU内核优化研究提供了标准化平台。研究团队同时开源了评估工具和基线结果,支持社区贡献。未来,随着TPU在AI训练和推理中的广泛应用,自动内核优化将有助于降低开发成本、提升硬件利用率。JAXBench的出现,有望像KernelBench推动GPU优化那样,加速TPU领域的技术迭代,使更多研究者和工程师能够参与这一关键环节的改进。

# TPU # 内核优化 # 基准测试 # JAX # Pallas

来源:Heooo AI工具导航