AI 推理术语表
软件栈

TorchAX

也称为 TorchAX、torchax、tpu-inference 后端、PyTorch 到 JAX 翻译层

先用大白话

TorchAX 通过把每个操作悄悄翻译成 JAX,让 PyTorch 模型代码能在 TPU 上运行,这种方式正被 TorchTPU 取代。

技术定义

TorchAX 是一个翻译层,通过 __torch_dispatch__ 钩子拦截 PyTorch 的 ATen 操作并以 JAX 操作执行,每个 torchax.tensor.Tensor 背后都是一个 jax.Array。

工程细节

vLLM 的 TPU 支持经历了三个阶段。最初的原型使用 PyTorch/XLA 的惰性执行,把操作收集成图交给 XLA。当前公开的 tpu-inference 后端在有 TPU 优化的 JAX 模型时直接使用,否则通过 TorchAX 运行 PyTorch 模型。权重和 KV cache 等状态通过 functionalization 显式暴露给 JAX,让 jax.jit 能捕获每一步,vLLM 因此绕过了自己常规的 torch.compile 路径,因为 JAX 流水线已经负责图捕获。SGLang-JAX 是一个独立的 JAX 原生引擎,做了类似取舍。Google 当时选择 JAX,是因为它有更成熟的 TPU 原语和并行支持。

为什么重要

跨两个框架翻译在底层优化、paged attention 以及让 vLLM 的 worker 模型适配 TPU 执行方面反复出问题,还迫使引擎功能在 JAX 侧重新实现。这些代价促使 Google、PyTorch、vLLM 和 SGLang 转而构建 TorchTPU。

如何在 InferenceX 中解读

在 TorchAX 下开发的 Pallas kernel,包括 InferenceX Official Preview 中描述的 MoE、GDN 和 paged attention 工作,都可以延续到 TorchTPU,因为两套栈都能调用 Pallas 和 JAX kernel。TorchTPU 发布后,TorchAX 本身将被弃用。