Google MaxKernel 自动生成 TPU 核,五十项几何均值加速 1.58 倍

MaxKernel: Agentic Kernel Generation for TPUs

Shangkun Wang, Nina Cai, Charles Hoong, Julian Walker, Gerson Kroiz, George Vanica, Deepak Patil, Andi Gavrilescu, Hassan Sipra, Sethu Sankaran

cs.AI, cs.PF, cs.PL

2026-09-04

Google 提出 MaxKernel,把规划、Pallas 实现、编译修复、冻结测试、自动调参和 XProf 剖析串成闭环,自动写 TPU 核。JaxBench 五十题并行搜索几何均值加速 1.58 倍;八个生产核 2.32 倍,超过人手写的 2.02 倍。

这篇在解决什么

写高性能 TPU 核,工程上一直卡在同一处。XLA 这种通用编译器吃不下注意力变体、稀疏算子和高度融合的算子,工程师只好绕过它,用 JAX 的 Pallas 手写。Pallas 要求人自己管 HBM 和 VMEM 两级存储、DMA 流水、多维 tiling。CUDA 和 Triton 在 GPU 上已经够难,TPU 的 API 更硬,报错更不透明。

零样本让大模型直接吐代码,过不了编译这一关。JaxBench 五十题上,100 次独立采样再取最快正确解,编译和数值正确都只有 10/50,几何均值加速 1.08 倍,几乎贴着 XLA 基线。代码能编过、数值对,还不等于快;要快,必须对着真实硬件的 trace 改 tiling 和访存。Halide、TVM、Ansor 把算法和调度拆开,仍然要领域知识或昂贵搜索。Google 这篇把编译器反馈、XProf 剖析、冻结测试集和搜索图绑成一套多 Agent 系统,叫 MaxKernel。

方法

共享的子 Agent 负责规划、把规划写成 Pallas、根据编译错误修、合成测试、在真机上跑、自动搜 block 和 tile、用 XProf 抽延迟、带宽和计算密度。知识库走 RAG,只收 Pallas、Mosaic、XLA 的文档和性能手册,明确不放人手写核,避免检索直接抄到专家实现。

三种编排共用这套子 Agent。

HITL 一次只跑一个子 Agent,停下来等人看规划稿或代码再往下。论文没给这条路径的量化数字。

Auto 把测试集先根据参考实现合成并冻结,堵住实现 Agent 改测试来混过正确性。随后循环「规划 → 实现 → 编译 → 测试 → 调参 → Profiling」,剖析结果喂回下一轮规划。编译修不出来或数值对不上,就短路回规划。全程留快照,结束时回滚到延迟最低的合法版本。单轨迹限 5 轮,每题跑 5 次报中位数。

图搜索把每次 Auto 产出当成图上一个节点,节点里装着代码、规划和实测指标,一次展开开一个新会话,避免上下文撑爆。并行搜索是 5 条互不剪枝的 Auto 轨迹,每条 5 轮,取最快正确解。Beam 宽 3、深 3、每节点 2 个分支,内层只给 2 轮,靠剪枝换广度。

全部实验用 Gemini 3.1 Pro,硬件是 TPU v6e。正确性用 jnp.allclose,多数 atol 和 rtol 取 1e-2,部分 bf16 任务放到 1e-1。延迟只计片上时间,扣掉 host 侧 dispatch 和编译。几何均值把慢于 XLA 的结果地板设成 1.0 倍。

结果

JaxBench 50 题:17 个来自常见 LLM 算子,33 个从 KernelBench 改编的融合算子。

方法编译正确几何均值加速fast1
Best-of-100 零样本10/5010/501.086/50
Auto 中位数49/5048/501.3922/50
并行搜索50/5050/501.5834/50
Beam50/5050/501.4931/50

Auto 的加速区间是 [1.19, 1.42],单轨迹会卡在局部可行但不够快的编译态。并行搜索把 fast1 拉到 34/50,加速阈值 p=2.0 时仍有 24% 的题过线。

八个有人手写 Pallas 的生产核,相对 XLA 的几何均值(地板 1×):人手 2.02 倍,并行搜索 2.32 倍,Beam 1.78 倍。八个里七个 Agent 更快。Paged Attention 上并行搜索 6.74 倍,人手 2.41 倍;Sparse Attention 是 5.03 对 2.45。翻车的是 Ragged Paged Attention:人手 4.65 倍,Agent 只有 1.42 倍。MLA 上人手写出 0.69 倍,比 XLA 还慢,Agent 做到 1.21–1.23 倍。GEMM 三方都在 1.02–1.03 倍,这块 XLA 已经够好。

真实模型核相对 JAX 或已有 Pallas:

负载基线MaxKernel对比
MLA v1 延迟3.127 ms(Pallas)2.856 ms降 8.68%
Qwen3-Next GDN 前向17.09 ms(JAX)10.47 ms1.63×
Qwen3-Next 训练步84.51 ms(JAX)17.99 ms4.70×
DeepSeek-V4 Sparse prefill27.08 ms(JAX)3.449 ms7.85×
DeepSeek-V4 Sparse decode0.047 ms(JAX)0.02 ms2.36×
Mamba v2 SSD0.211 ms(JAX)0.192 ms1.10×
StreamIndextopk0.415 ms(JAX)0.250 ms1.66×

MLA v1 吞吐从 116.7 提到 127.8 TFLOPS。系统还修过 Ragged Page Attention v3 prefill 的崩溃:自动加上 ALU clamp,处理左填充,避免负 slice 和脏 DMA。

为什么重要

给已经在 TPU 上写 Pallas 的人,这是一套能开源跑的闭环:参考实现进,编译错误和 XProf trace 回灌,出一份能过数值门的核。JaxBench 上从「100 次盲采样五分之一能编过」到「并行搜索 50/50 正确、几何 1.58 倍」,增量主要来自硬件反馈,不是更大的基座模型。

八个生产核几何均值超过人手,不等于处处能替人。Ragged Paged Attention 人手仍快三倍以上,GEMM 几乎没空间。并行搜索的 1.58 倍是 5 条轨迹取最优,计算量大约是单次 Auto 的五倍;论文没报 token 和墙钟。HITL 只给了交互设计,没有对照数字。

能用的场景很具体:已有 JAX 参考、有 TPU v6e、愿意烧多轨迹搜索。GPU 上的 CUDA/Triton、Trainium 的 NKI、MTIA 只在引言里被点名,实验全在 Pallas。

局限与存疑

作者自己写了后续:换进化或贪心搜索、让知识库随实验增长、用有状态的混合策略处理探索和利用。模型只测了 Gemini 3.1 Pro。

几何均值地板 1.0 倍会抬高总分,慢于 XLA 的核在均值里被当成「没变慢」。数值门从 1e-2 放到 1e-1,对训练向后传的核偏松。Paged Attention 的 6.74 倍在表里属于并行搜索,正文有一处写成 Beam,以表为准。

并行搜索相对 Beam 的优势,和「每条轨迹 5 轮对 2 轮」绑在一起,不是纯粹的算法对比。知识库故意不放手写核,评测干净,落地时这个选择可能反过来。论文没和 AutoComp 等先前 LLM 核优化工作做同任务对照,也没报搜索成本。

术语

原文与代码

相关论文

全部论文解读