跳到正文
原文
LMSYS:Blog(Chatbot Arena 团队)·· 2 小时前精选AI 评分63

SGLang 团队发布开源 TPU 原生推理引擎 SGLang-Jax

Blog SGLang-Jax: An Open-Source Solution for Native TPU Inference We're excited to introduce SGLang-Jax, a state-of-the-art open-source inference engine built entirely on Jax and XLA. It leverages SGLang's high-performance server architecture and uses Jax to compile... The SGLang-Jax Team October 29, 2025

AI 导读

LMSYS Chatbot Arena 团队发布开源推理引擎 SGLang-Jax,完全基于 Jax 和 XLA 构建,提供原生 TPU 推理,基准显示其达到或超过其他 TPU 推理方案。

推荐理由

官方团队介绍了基于 Jax 的 TPU 推理引擎的关键优化和基准结果,读者可以据此评估它在现有推理方案中的位置。

正文 · AI 翻译

我们很高兴推出 SGLang-Jax,这是一个完全基于 Jax 和 XLA 构建的最先进的开源推理引擎。 它利用 SGLang 的高性能服务器架构,并使用 Jax 编译模型的前向传播。 通过结合 SGLang 和 Jax,该项目实现了快速的原生 TPU 推理,同时保持对连续批处理、前缀缓存、张量和专家并行、投机解码、内核融合以及高度优化的 TPU 内核等高级功能的支持。

基准测试表明,SGLang-Jax 与其他 TPU 推理解决方案持平或更优。 源代码可在 https://github.com/sgl-project/sglang-jax 获取。

为什么选择 Jax 后端?

虽然 SGLang 最初构建于 PyTorch 之上,但社区一直渴望对 Jax 的支持。
我们构建 Jax 后端有几个关键原因:

  • Jax 从底层开始就是为 TPU 设计的。为了在不妥协的情况下实现最大性能,Jax 是明确的选择。随着 Google 扩大 TPU 的公共访问,我们预计 Jax + TPU 将获得显著的发展势头,并实现高性价比的推理。
  • 领先的 AI 实验室——包括 Google DeepMind、xAI、Anthropic 和 Apple——已经依赖 Jax。在训练和推理中使用相同的框架可以减少维护开销,并消除两个阶段之间的偏差。
  • Jax + XLA 是一个经过验证的、编译驱动的技术栈,在 TPU 上表现出色,并且在各种定制的类 TPU AI 芯片上表现良好。

架构

下图展示了 SGLang-Jax 的架构。整个技术栈是纯 Jax,因此代码简洁、依赖极少。

在输入侧,它通过 OpenAI 兼容的 API 接受请求,并利用 SGLang 高效的前缀缓存 RadixCache 以及其重叠调度器实现低开销批处理。 调度器为不同的批大小预编译 Jax 计算图。 在模型侧,我们使用 Flax 实现模型,并使用 shard_map 实现各种并行策略。 两个核心算子——注意力和 MoE——实现为自定义 Pallas 内核。

SGLang-Jax 的架构

关键优化

集成 Ragged Paged Attention v3

我们集成了 Ragged Paged Attention V3(RPA v3)并扩展它以支持 SGLang 功能:

  • 我们根据不同的场景调整内核网格块配置,以实现更好的性能。
  • 我们使其与 RadixCache 兼容。
  • 为了支持 EAGLE 投机解码,我们为 RPA v3 添加了自定义掩码,用于验证阶段。

减少调度开销

前向传播过程中 CPU 和 TPU 上的顺序操作可能会影响性能。然而,不同设备上的操作可以解耦——例如,在 TPU 上启动计算并立即准备下一批要运行的批次。为了提高性能,我们的调度器将 CPU 处理与 TPU 计算重叠。

在重叠事件循环中,调度器使用结果队列和线程事件来流水线化 CPU 和 TPU 工作。当 TPU 处理批次 N 时,CPU 准备批次 N+1。为了最大化 CPU 和 TPU 之间的重叠,SGLang-jax 根据性能分析结果仔细安排操作顺序。对于 Qwen/Qwen3-32B,我们将预填充和解码之间的时间间隔从约 12ms 减少到 38us,从约 7ms 减少到 24us。更多细节可以在我们之前的博客中找到。

使用重叠调度器的性能分析。批次之间的间隔极小。

不使用重叠调度器的性能分析。注意批次之间存在较大的间隔(CPU 开销)。

MoE 内核优化

MoE 层目前支持两种实现策略:EPMoE 和 FusedMoE。 在 EPMoE 中,我们集成了 Megablox GMM 算子,取代了此前基于 jax ragged_dot 的实现。 Megablox GMM 专为 MoE 工作负载设计,能够高效处理由 group_sizes 描述的可变大小专家分组,消除不必要的计算和非连续内存访问。在典型配置下,该算子相比 jax 原生的 ragged_dot 实现可带来 3–4 倍的端到端(e2e)ITL 加速。 结合高效的 token 重排(permute/unpermute)、通过 ragged_all_to_all 实现的专家并行通信,以及自适应分块策略,EPMoE 显著提升了整体吞吐量,并且在需要跨设备并行且专家数量较多的场景中表现良好。 相比之下,FusedMoE 使用稠密 einsum 运算融合所有专家计算,没有跨设备通信开销。它更适合单个专家较大但专家总数较少的情况(例如 < 64 个专家)。它也可作为轻量级回退方案,便于调试和正确性验证。

投机解码

SGLang-jax 实现了基于 EAGLE 的投机解码,也称为多 token 预测(MTP)。 这种先进的投机解码技术通过使用轻量级草稿头预测多个 token 来加速生成,随后在完整模型的一次前向传播中并行验证这些 token。 为了实现基于树的 MTP-Verify,SGLang-jax 在 Ragged Paged Attention V3 之上增加了非因果掩码支持,从而在验证阶段能够并行解码基于树的非因果草稿 token。 我们目前支持 Eagle2 和 Eagle3,并计划继续优化内核实现,并在 MTP 的各个阶段增加对不同注意力后端的支持。

TPU 性能

经过上述所有优化后,SGLang-Jax 达到或超越了其他 TPU 推理方案。 与 GPU 方案相比,TPU 上的 SGLang-Jax 同样具有竞争力。

你可以在 https://github.com/sgl-project/sglang-jax/issues/297 找到完整的基准测试结果和说明。

使用方法

安装 SGLang-Jax 并启动服务器

安装:

# with uv
uv venv --python 3.12 && source .venv/bin/activate
uv pip install sglang-jax

# from source
git clone https://github.com/sgl-project/sglang-jax
cd sglang-jax
uv venv --python 3.12 && source .venv/bin/activate
uv pip install -e python/

启动服务器:

MODEL_NAME="Qwen/Qwen3-8B"  # or "Qwen/Qwen3-32B"

jax_COMPILATION_CACHE_DIR=/tmp/jit_cache \
uv run python -u -m sgl_jax.launch_server \
--model-path ${MODEL_NAME} \
--trust-remote-code \
--tp-size=4 \
--device=tpu \
--mem-fraction-static=0.8 \
--chunked-prefill-size=2048 \
--download-dir=/tmp \
--dtype=bfloat16 \
--max-running-requests 256 \
--page-size=128

通过 GCP 控制台使用 TPU

你可以在控制台的 Menu → Compute Engine 下找到 TPU 选项,然后点击 Create TPU。 注意:只有特定区域支持特定的 TPU 版本。请记得将 TPU 软件版本设置为 v2-alpha-tpuv6e。 在 Compute Engine 菜单下,进入 Settings → Metadata,点击 SSH Keys 按钮,并添加你的公钥。 TPU 服务器创建完成后,你可以使用控制台中显示的外部 IP 和公钥用户名登录。 另请参阅:https://docs.cloud.google.com/tpu/docs/setup-gcp-account

通过 SkyPilot 使用 TPU

我们推荐使用 SkyPilot 进行日常开发。 你可以在 sglang-jax 仓库中快速设置 SkyPilot,并找到用于启动开发机器和运行测试的脚本。

为 GCP 安装 SkyPilot:https://docs.skypilot.co/en/latest/getting-started/installation.html#gcp 然后启动 sgl-jax.sky.yaml:

sky launch sgl-jax.sky.yaml --cluster=sgl-jax-skypilot-v6e-4 --infra=gcp -i 30 --down -y --use-spot

该命令将跨区域查找成本最低的 TPU 抢占式实例,并在空闲 30 分钟后自动关闭该实例。它还会为你安装 sglang-jax 环境。 设置完成后,你可以直接使用 ssh cluster_name 登录,而无需跟踪外部 IP 地址。

路线图

社区正在与 Google Cloud 团队及多个合作伙伴共同推进以下路线图。

  • Model support and optimizations
    • 优化 Grok2、Ling/Ring、DeepSeek V3 和 GPT-OSS
    • 支持 MiMo-Audio、Wan 2.1、Qwen3 VL
  • TPU-optimized kernels
    • 量化内核
    • 通信与计算重叠内核
    • MLA 内核
  • RL integration with tunix
    • 权重同步
    • Pathways 与多主机支持
  • Advanced serving features
    • 预填充-解码分离
    • 分层 KV 缓存
    • 多 LoRA 批处理

致谢

SGLang-jax 团队:sii-xinglong、jimoosciuc、Prayer、aolemila、JamesBrianD、zkkython、neo、leos、pathfinder-pf、Jiacheng Yang、Hongzhen Chen、Ying Sheng、Ke Bao、Qinghan Chen

Google:Chris Yang、Shun Wang、Michael Zhang、Xiang Li、Xueqi Liu

InclusionAI:Junping Zhao、Guowei Wang、Yuhong Guo、Zhenxuan Pan

来源:LMSYS:Blog(Chatbot Arena 团队) · lmsys.org