编译器优先状态空间二元性与面向推理的O(1)自回归缓存
机器学习
2026-03-11 v1 人工智能
分布式、并行与集群计算
性能
摘要
状态空间模型版本通常与融合的CUDA和Triton内核捆绑,继承了对NVIDIA硬件的严格依赖。我们展示,Mamba-2的状态空间二元性算法——对角线状态结构、可块化递归以及以einm形式计算的静态控制流——可顺利映射到XLA的融合和平铺 passes 所优化的结构,使得自定义内核可选而非必需。我们将完整的推理路径(预填充、缓存自回归解码)实现为标准原语,无需手写内核,实现了体系结构的理论O(1)状态管理作为编译时设备缓存,生成期间无需主机同步。该实现在CPU、NVIDIA GPU和Google Cloud TPU上均可 unmodified运行,源于单一的JAX源。在TPU v6e上进行五个模型规模(130M至2.7B参数)测试中,XLA生成代码在单流预填充上达到约140 TFLOPS(15% MFU),解码时达到最高64%带宽利用率。贪心解码在64步内与PyTorch/CUDA参考实现逐token匹配,隐藏状态误差在float32舍入容限内。该模式可移植至满足相同结构条件的任何SSM recurrence,在任何拥有成熟XLA后端的平台上。该实现已公开发布于https://github.com/CosmoNaught/mamba2-jax,并合并至Bonsai JAX模型库。
引用
@article{arxiv.2603.09555,
title = {Compiler-First State Space Duality and Portable $O(1)$ Autoregressive Caching for Inference},
author = {Cosmo Santoni},
journal= {arXiv preprint arXiv:2603.09555},
year = {2026}
}
备注
18 pages, 6 figures. Code available at: https://github.com/CosmoNaught/mamba2-jax