vLLM TPU:在 TPU 上支持 PyTorch 和 JAX 的全新统一后端
vLLM TPU 现在由 tpu-inference 提供动力,这是一个表现力强且功能强大的新型硬件插件,将 JAX 和 PyTorch 统一在单个降低路径(lowering path)下。它不仅比上一代 vLLM TPU 更快,还提供了更广泛的模型覆盖和功能支持。vLLM TPU 是一个供开发者使用的框架,旨在:
- 在开源领域挑战 TPU 硬件性能的极限。
- 通过在 TPU 上高效运行 PyTorch 模型定义而无需任何额外的代码更改,同时将原生支持扩展到 JAX,从而为 JAX 和 PyTorch 用户提供更多的灵活性。
- 保留 vLLM 标准化:保持相同的用户体验、遥测和接口。
vLLM TPU
2025 年 2 月,正当 vLLM 的 V1 集成 初具规模时,一个由 Google 员工和 vLLM 核心贡献者组成的“小而精”的团队设定了一个目标:赶在 Cloud Next 2025 之前,在少数模型上发布高性能的 TPU 后端。在接下来的两个月里,他们遇到了几个挑战,即:
- vLLM V1 集成:团队必须集成到新的 V1 代码路径中,这需要一个新的不规则分页注意力内核(RPA v2)。这主要是为了支持分块预填充(chunked prefill)和前缀缓存(prefix caching)等功能。虽然这些 KV 缓存管理技术在 TPU 上很常见,但要以“TPU 友好”的方式结合 vLLM 的分页注意力进行设计极具挑战性。
- 多程序多数据 (MPMD):当时,vLLM 专门使用 MPMD 来协调跨进程通信。这与 TPU 以编译器为中心的编程模型形成鲜明对比,后者严重依赖单程序多数据 (SPMD) 来实现多设备和多主机通信的重叠。
- PyTorch/XLA (PTXLA):尽管使用 PyTorch/XLA 框架使得集成到 vLLM 变得更容易(因为它能够在 TPU 上原生运行 PyTorch 代码),但团队在优化技术栈的较低层级时遇到了许多挑战。
尽管存在这些障碍,团队仍将 Llama 3.1-8B 在 v6e-1 上的吞吐量性能提高了 3.6 倍,将 Llama 3.1-70B 在 v6e-8 上的吞吐量性能提高了 2.1 倍。vLLM TPU 也登上了 Cloud Next 的大舞台。您可以在此处查看这些工作负载的性能演进。
由 TPU-inference 驱动的 vLLM TPU
尽管采用 PTXLA 的 vLLM TPU 是一项重大成就,但我们需要继续在开源领域突破 TPU 性能的极限。我们还希望通过以最高效的方式在 TPU 上原生支持 PyTorch 和 JAX 模型,将 TPU 和 vLLM 生态系统聚集在一起。
PyTorch 和 JAX 的统一后端
这次使用 tpu-inference 进行的 vLLM TPU 重新设计旨在通过在单个统一的 JAX→XLA 降低路径中支持 PyTorch(通过 Torchax)和 JAX,来优化性能和扩展性。
与 PyTorch/XLA 相比,JAX 是一个更成熟的技术栈,通常为其原语提供卓越的覆盖范围和性能,特别是在实现复杂的并行策略时。
因此,vLLM TPU 现在使用 JAX 作为所有 vLLM 模型的降低路径,即使模型定义是用 PyTorch 编写的,也能从显着的性能提升中受益。这一决定使我们能够更快、更聪明地移动,抽象掉高级框架,专注于算子(kernel)开发和编译器优化。请记住,对于 XLA,Torchax 和 JAX 在编译前使用相同的高性能原语。您可以在此处阅读更多相关信息。
虽然这是我们目前的设计,但我们将始终致力于在 TPU 上实现最佳性能,并计划在未来为 vLLM TPU 评估在 TPU 上的原生 PyTorch 移植。
重要提示
要点 #1:vLLM TPU 现在使用 JAX 降低所有模型。在不对模型代码(例如 llama.py)进行任何更改的情况下,vLLM TPU 现在实现了约 20% 的更高吞吐量性能,这仅仅是因为它现在利用 JAX 成熟的高性能原语来生成随后由 XLA 编译的 HLO 图。
深入了解
-
安装
pip install vllm-tpu # a single install path因为 Torchax 和 JAX 在底层本质上只是 JAX,所以无论模型代码是用 PyTorch 还是 JAX 编写的,我们都可以利用相同的安装路径。这确保了依赖项保持一致,用户不必担心为不同模型管理不同的需求。
-
提供模型服务
MODEL_ID="google/gemma3-27b-it" # model registered in tpu-inference or vllm vllm serve $MODEL_ID在 TPU 上提供模型服务时,有两个模型注册表可供拉取模型代码:
让我们深入了解一下底层发生了什么:
这种统一工作通过利用 vLLM 社区的现有成果减少了重复工作,留出更多时间来优化 TPU 算子和 XLA 编译器。对于 PyTorch(通过 Torchax)和 JAX 模型,所有算子和编译器都是共享的。
重要提示
要点 #2:vLLM TPU 现在将默认运行 tpu-inference 中经过 TPU 优化的模型代码(如果存在);否则,它将回退到 vLLM 上游的 PyTorch 模型代码(通过 Torchax 使用 JAX 降低)。对于大多数用户来说,这只是一个实现细节。
如果 Torchax 可以在 TPU 上开箱即用地运行 PyTorch 模型代码,但仍然使用 JAX JIT 进行编译,为什么我们还要在 tpu-inference 中重写一些模型?这不是重复的吗?
我们为开发者提供了一些参考模型,以减少他们在开始为 TPU 优化模型之前的学习曲线(参见此处)。有趣的是,我们观察到 torchax 降低的模型和原生重新实现的 JAX 模型具有大致相同的性能,这证明了 torchax 在转换高级模型方面的效率。
真正的性能优势以及我们支持重新实现模型的原因,来自于针对 TPU 优化 JAX 代码并直接利用 TPU 架构的优势。
我们需要这种灵活性,是因为 vLLM 开发者在实现模型时的逻辑设计选择并不总是偏向于 TPU。这使得它们变得不同,不是因为 JAX 与 Torchax 的区别,而是因为 GPU 与 TPU 不同,需要不同的优化策略。
重要提示
要点 #3:对于任何模型,底层全是 JAX!除非实现中的逻辑差异导致 TPU 性能下降,否则模型可能不会从原生 JAX 重写中受益。尽管如此,保留重新实现模型的灵活性很重要,如果这意味着我们可以发挥出 TPU 的最佳性能。
Ragged Paged Attention V3:开源领域面向 TPU 推理的最灵活、高性能的注意力内核
虽然 Ragged Paged Attention v2 内核在性能上有了重大提升,但为了开箱即用地支持更多模型和用例,它需要变得更加灵活。
- RPA v2 只能支持 head dim 为 128 的模型规格。
- 更多模型:RPA v3 更加灵活,支持任意模型规格、量化数据类型和任意张量并行 (TP),开箱即可解锁更多模型。
- 由于顺序执行 KV 缓存更新和注意力操作,RPA v2 存在流水线效率低下的问题。
- 更好的性能:RPA v3 通过将 KV 缓存更新 (scatter) 融合到 RPA 内核中来提高流水线效率。这种设计现在在内核执行期间完全隐藏了 scatter 延迟。
- 在解码密集型或不同长度的预填充任务期间,RPA v2 可能会产生巨大的浪费。
- 改进的部署灵活性:RPA v3 将编译为 3 个子内核,解锁对仅预填充(prefill-only)、仅解码(decode-only)和混合批处理的支持。这种设计通过在运行时将正确的子内核与相应的请求匹配,显着节省了直接内存访问 (DMA) 和计算。
- 这还带来了额外的好处,即解锁了更复杂的部署模式,如解耦服务(disaggregated serving)。
- 虽然 RPA v2 与第一个 TPU 原型相比实现了显著的吞吐量改进,但它缺乏灵活性。
- 不妥协: RPA v3 不会为了灵活性而牺牲性能,事实上,它在 Trillium (v6e) 上的吞吐量比 RPA v2 提高了约 10%。模型现在也可以在 v5p 上运行(尽管需要额外的调优)。
我们很快将撰写关于 RPA v3 的技术深度解析,请关注我们的文档。
重要提示
要点 #4:RPA v3 既灵活又高效,是开源领域生产级 Pallas 内核开发的优秀参考。我们很高兴看到 TPU 友好的 MoE 和 MLA 内核很快也会以类似的方式在开源社区落地。
单程序多数据 (SPMD)
此版本引入了单程序多数据 (SPMD) 作为 vLLM TPU 的默认编程模型。与之前的多工作器(multi-worker)模型(改编自 GPU 范式)不同,SPMD 是 XLA 编译器的原生模式。开发者为一个单一的庞大设备编写代码,XLA 编译器会自动划分模型和张量,插入通信操作以实现最佳执行。
重要提示
要点 #5:SPMD 支持高级优化,例如通信与计算的重叠。SPMD 代表了向更深层次原生 TPU 集成的战略转型,承诺通过以 TPU 为中心、编译器优先的运行模式提供更高性能。
总结
|
|
vLLM TPU 自 2025 年 2 月的原理样机性能以来已经走过了漫长的道路,在相同的工作负载下达到了近 2x-5x 的性能,同时还改进了模型覆盖范围和易用性。
重要提示
要点 #6:今天,vLLM TPU 的性能几乎是 2025 年 2 月第一个 TPU 原型的 5 倍。有了这个新的基础,开发者和研究人员现在将能够在开源领域以前所未有的速度突破 TPU 推理性能的界限。
模型、功能及后续计划
我们可以将此版本视为基础版本,因为 vLLM TPU 现在将在开源领域定期发布版本。随每个新版本的发布,CI/CD 将发布经过验证的 vLLM 原生模型文档表。我们还将维护一个经过压力测试的 tpu-inference 模型列表,主要作为 JAX 用户的参考。所有功能在发布前也将经过严格测试。
支持的模型系列
- 稠密模型 (Dense)
- 多模态模型(仅限 tpu-inference 模型)
说明
关于模型支持的说明:在落地更多功能之前,我们建议从此处经过压力测试的模型列表开始。我们仍在 tpu-inference 中落地各种组件,这些组件将提高大规模、高复杂度模型(XL MoE、+视觉编码器、MLA 等)的性能。如果您希望我们优先处理某些特定内容,请在此处提交 GitHub 功能请求。
支持/已验证的 TPU 代际
- Trillium (v6e), v5e
功能特性
- 前缀缓存 (Prefix caching)
- 分块预填充 (Chunked Prefill)
- 多模态输入
- 单程序多数据 (SPMD)
- 结构化解码
- 投机解码:Ngram
- 树外(Out-of-tree)模型支持
- 优化的运行时采样(top k, top p, temperature, logit 输出)
- 量化(权重、激活和 KV 缓存)
TPU 友好算子
- Ragged Paged Attention V3
- 集合通信 Matmul
- 量化 Matmul、Attention 和 KV 缓存
实验性功能
- v5p
- 多模态(通过 Torchax)
- Multi-lora
- 投机解码:基于树的 Eagle 3
- 单主机 P/D 解耦服务
未来计划
- Sparsecore 卸载
- 投机解码:Eagle 3, MTP
- TPU 友好算子
- XL MoE
- MLA
- 集成
- 欢迎贡献!
立即体验!
您可以在 Google Cloud 上进行尝试,包括 Google Kubernetes Engine (GKE)、Compute Engine 和 Vertex AI。有关安装说明和开发者指南,请查看以下资源:
Google Cloud 教程:GKE 请看此处,Vertex AI 请看此处
致谢
我们想对 vLLM 社区在这项工作中的持续支持表示最诚挚的谢意。特别感谢 Woosuk Kwon 牵头 TPU 的 V0 实现并继续支持我们不断壮大的团队。我们还要特别鸣谢 Simon Mo, Robert Shaw, Michael Goin, Yanping Huang 在整个工作过程中的宝贵指导。同时也要特别感谢 Nicolo Lucchesi, Alexander Matveev, Akshat Tripathi, 和 Saheli Bhattacharjee,他们是 V1 集成和 Cloud Next 冲刺中不可或缺的一部分。