【第727期】JAXBench:自主TPU内核优化评测基准Seventy3

【第727期】JAXBench:自主TPU内核优化评测基准

25分钟 ·
播放数5
·
评论数0

今天的这篇的主题是: JAXBench: Benchmarking Autonomous TPU Kernel Optimization

Seventy3 是一档借助 NotebookLM 解读前沿论文的播客——不是念摘要,也不是泛泛而谈,而是让 AI 帮你把论文掰开揉碎,说人话。我们蹲在人工智能、大模型、机器人算法、crypto 这几个领域,每期挑一篇值得关注的工作,用对话的方式聊给你听。你可以在通勤路上听,做实验的时候听,也可以当背景音放着,让新知识自然地长进脑子里。

如果你正在做有意思的研究,想让更多人看到——不管你是博士生、独立研究者还是实验室的博士后——把论文发给我们,我们帮你用 AI 做一期深度解读,让你的工作被更多同路人发现。

联系小助手微信:eccstartup(加群/投稿论文)

Summary

严格的基准测试通过确立可供攀登的共同目标,推动了自主 GPU 内核性能优化领域的进步,但 TPU 领域尚无同等效力的基准。我们推出了 JAXBench,这是一个在 Google Cloud TPU 上用于 AI 生成内核优化的 TPU 原生基准测试套件。

JAXBench 包含 50 个既具备实际相关性又留有优化空间(headroom)的 JAX 工作负载。我们从公共 MaxText 库中的架构(如 Llama-3.1、DeepSeek-V3、Mixtral、Mamba-2 和 AlphaFold2)中提取了 17 个生产级 ML 算子,并从 KernelBench 中翻译了 33 个经正确性验证的算子,同时设置了能够实现高 TPU v6e MXU 利用率的新问题规模。在 17 个生产级算子中,有 8 个附带了来自公共 Tokamax 库且经过分块大小调优(block-size tuned)的手工优化 Pallas 内核,以确立专家级上限基线。

我们评估了四种基于反馈的方法在为 JAXBench 生成候选 Pallas 内核上的表现。在配合 Gemini 3 Flash 的全套基准测试中,我们发现对于像 Pallas 这样缺乏丰富文档的 DSL 而言,特定于目标的上下文比模型规模更为关键。提供精心挑选的 TPU 文档上下文,可将单样本正确率(per-sample correctness)从 5.8% 提升至 37.3%,并在 50 个基准测试中解决了 48 个,实现了 1.28 倍的几何平均加速比。

一旦实现正确性,搜索结构就会带来显著收益:Autocomp 的束搜索(beam-search)流水线相比 XLA 达到了 1.36 倍的几何平均加速比。在 8 个手工调优的内核上,Autocomp 相比 XLA 取得了 1.60 倍的几何平均加速比,恢复了 Tokamax 2.08 倍上限的大部分性能,但在专用的分页(paged)和不规则(ragged)注意力算子方面仍稍显逊色。

高质量的 TPU 内核优化仍是一项极具挑战性的任务,我们现发布 JAXBench 基准测试、评估框架(harness)以及基线结果,以支持开源社区的贡献。

原文链接:arxiv.org