SheepNav
精选今天0 投票

JAXBench:为TPU内核自动优化设立新标杆

GPU领域的自动内核优化得益于像KernelBench这样的基准测试套件,它们为算法提供了一个共同的优化目标。然而,TPU(张量处理单元)领域却长期缺乏类似的标准化评估工具。近日,一篇发表于arXiv的论文(编号2607.20466)介绍了JAXBench——一个专为Google Cloud TPU设计的AI生成内核优化基准套件,旨在填补这一空白。

JAXBench包含50个JAX工作负载,这些负载既具有实际相关性,又为优化留出了充分的改进空间。其中,17个生产级ML算子直接提取自MaxText公共库中的主流架构,包括Llama-3.1、DeepSeek-V3、Mixtral、Mamba-2和AlphaFold2。另外33个算子来自KernelBench的移植版本,经过正确性验证并重新设定了问题规模,以确保在TPU v6e上实现高MXU利用率。此外,8个生产算子还附带了来自Tokamax公共库的手工优化Pallas内核,并经过块大小调优,作为专家级的上限基线。

论文评估了四种基于反馈的方法,用于为JAXBench生成候选Pallas内核。实验结果表明,在像Pallas这样文档稀疏的DSL(领域特定语言)中,目标特定的上下文比模型规模更重要。使用Gemini 3 Flash模型时,通过引入经过整理的TPU文档,单样本正确率从5.8%提升至37.3%,成功解决了50个基准中的48个,并实现了1.28倍的几何平均加速。一旦正确性得到保证,搜索结构能带来显著收益。例如,Autocomp的束搜索流水线相比XLA实现了1.36倍的几何平均加速。在8个手工调优内核上,Autocomp达到1.60倍的几何平均加速,接近Tokamax提供的2.08倍上限,但在专门的分页和稀疏注意力算子方面仍有差距。

JAXBench的发布为TPU内核优化社区提供了一个标准化的评估平台。研究团队公开了基准套件、评估工具和基线结果,以促进开源贡献。这一工作不仅展示了自动优化在TPU上的潜力,也揭示了当前方法的局限性——尤其是在处理复杂注意力机制时。随着TPU在AI训练和推理中的广泛应用,JAXBench有望成为推动硬件与软件协同设计的关键工具。

延伸阅读

  1. DecodeShare:追踪LLM解码时刻共享子空间的新方法
  2. InferenceBench:AI智能体在开放式大模型推理优化中的真正能力测试
  3. DC-Leap:无需训练的dLLM加速新方法,最高提速105倍
查看原文