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/50 | 10/50 | 1.08 | 6/50 |
| Auto 中位数 | 49/50 | 48/50 | 1.39 | 22/50 |
| 并行搜索 | 50/50 | 50/50 | 1.58 | 34/50 |
| Beam | 50/50 | 50/50 | 1.49 | 31/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 ms | 1.63× |
| Qwen3-Next 训练步 | 84.51 ms(JAX) | 17.99 ms | 4.70× |
| DeepSeek-V4 Sparse prefill | 27.08 ms(JAX) | 3.449 ms | 7.85× |
| DeepSeek-V4 Sparse decode | 0.047 ms(JAX) | 0.02 ms | 2.36× |
| Mamba v2 SSD | 0.211 ms(JAX) | 0.192 ms | 1.10× |
| StreamIndextopk | 0.415 ms(JAX) | 0.250 ms | 1.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 核优化工作做同任务对照,也没报搜索成本。