AI 推理术语表
软件栈

JAX

也称为 JAX、JAX 框架、jax.jit、Google JAX

先用大白话

JAX 是 Google 的 Python 数组计算框架,通过 XLA 编译函数,一直是编程 TPU 的原生方式。

技术定义

JAX 是一个 Python 库,在 NumPy 风格数组上提供 jit、grad、vmap 等可组合的函数变换,经 XLA 编译,是 TPU 的主要第一方框架。

工程细节

JAX 程序是函数式的:模型权重和 KV cache 等状态显式传入传出,jax.jit 才能把一步计算追踪成图。这种模型与 XLA 契合,让 JAX 成为 TPU 原语和并行支持最成熟的路径,也是第一个外部 vLLM TPU 后端通过 TorchAX 把 PyTorch 翻译成 JAX、以及 SGLang-JAX 作为独立 JAX 原生引擎存在的原因。TorchTPU 改变了这一关系:PyTorch 成为面向用户的框架,但底层仍可调用 JAX 自定义 kernel 和 Pallas,TPU-Sync 也原生支持 JAX 和 TorchTPU。

为什么重要

开源推理生态是用 PyTorch 写的。要求引擎跨越 PyTorch 到 JAX 的边界拖慢了 TPU 上的功能对齐和模型 bring-up,因此 Google 把推理栈迁向 PyTorch 原生,同时把 JAX 和 Pallas 保留给 kernel 和内部负载。

如何在 InferenceX 中解读

Qwen3.5 最初的 TPU 支持先添加了 causal Conv1D 和 Gated DeltaNet 的纯 JAX 实现,之后才由 Pallas kernel 优化。InferenceX Official Preview 的数据来自 TorchTPU vLLM 栈,而不是经 JAX 翻译的 tpu-inference 后端。