vLLM TPU:支持 TPU 上 PyTorch 和 JAX 的全新统一后端

阅读时间 11 分钟
Google 团队

vLLM TPU 现在由 tpu-inference 提供支持,这是一个富有表现力且强大的新型硬件插件,将 JAXPyTorch 统一在单一降级(lowering)路径下。它不仅比上一代 vLLM TPU 更快,而且提供了更广泛的模型覆盖和功能支持。vLLM TPU 是一个让开发者能够实现以下目标的框架:

  1. 在开源领域突破 TPU 硬件性能极限。
  2. 为 JAX 和 PyTorch 用户提供更多灵活性,无需任何额外代码更改即可在 TPU 上高效运行 PyTorch 模型定义,同时将原生支持扩展到 JAX。
  3. 保持 vLLM 标准化:维持相同的用户体验、遥测和接口。

vLLM TPU

2025 年 2 月,当 vLLM 的 V1 集成初具雏形时,一个由 Google 员工和 vLLM 核心贡献者组成的“小而精”的团队设定了一个目标:赶在 Cloud Next 2025 之前,在少量模型上推出高性能的 TPU 后端。在接下来的两个月里,他们遇到了几个挑战,具体如下:

  • vLLM V1 集成:团队必须集成到新的 V1 代码路径中,这需要一个新的不规则分页注意力内核(RPA v2)。这样做主要是为了支持分块预填充(chunked prefill)和前缀缓存(prefix caching)等功能。尽管这些 KV 缓存管理技术在 TPU 上很常见,但要将其与 vLLM 的分页注意力机制以一种“TPU 友好”的方式结合起来设计,极具挑战性。
  • 多程序多数据(MPMD:当时,vLLM 完全使用 MPMD 来协调跨进程通信。这与 TPU 以编译器为中心的编程模型形成了鲜明对比,后者严重依赖单程序多数据(SPMD)来进行重叠的多设备和多主机通信。
  • PyTorch/XLA (PTXLA):尽管使用 PyTorch/XLA 框架因为其能够在 TPU 上原生运行 PyTorch 代码的特性而简化了与 vLLM 的集成,但团队在栈的底层进行优化时还是遇到了一些挑战。

尽管存在这些障碍,团队还是将 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 编写的,也能从中获得显著的性能提升。这一决策使我们能够更快、更智能地行动,抽象掉更高级别的框架,专注于内核开发和编译器优化。请记住,对于 XLA 而言,Torchax 和 JAX 在编译前使用相同的高性能原语。您可以点击此处阅读更多相关信息。

虽然这是我们当前的设计,但我们将始终致力于在 TPU 上实现最佳性能,并计划在未来为 vLLM TPU 评估原生 PyTorch 移植版本。

重要:要点 #1:vLLM TPU 现在使用 JAX 对所有模型进行降级处理。无需对模型代码(例如 llama.py)进行任何更改,vLLM TPU 现在即可实现约 20% 的吞吐量提升,原因仅仅是它现在利用了 JAX 成熟的高性能原语来生成由 XLA 编译的 HLO 图。

深入了解

  1. 安装

    pip install vllm-tpu # a single install path

    由于 Torchax 和 JAX 本质上都是底层的 JAX,因此无论模型代码是用 PyTorch 还是 JAX 编写的,我们都可以利用相同的安装路径。这确保了依赖项的一致性,用户无需担心为不同模型管理不同的需求。

  2. 模型服务

    MODEL_ID="google/gemma3-27b-it" # model registered in tpu-inference or vllm 
    vllm serve $MODEL_ID

    在 TPU 上提供模型服务时,有两种模型注册表可供拉取模型代码:

    1. tpu-inference (默认,列表)
    2. vllm (在 vLLM 上游维护,列表)

让我们仔细看看底层发生了什么


这一统一工作通过利用 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 内核在性能上有了重大提升,但为了开箱即用支持更多的模型和用例,它需要变得更加灵活。

  1. RPA v2 仅支持头维度(head dim)为 128 的模型规格。
    • 更多模型:RPA v3 更加灵活,支持任意模型规格、量化数据类型以及任意张量并行(TP),实现了更多模型的开箱即用。
  2. RPA v2 由于顺序执行 KV 缓存更新和注意力操作,导致流水线效率低下。
    • 更好的性能:RPA v3 通过将 KV 缓存更新(分散操作)融合到 RPA 内核中,提高了流水线效率。这种设计现在在内核执行期间完全隐藏了分散操作的延迟。
  3. RPA v2 在解码密集型或不同长度的预填充任务中可能会产生严重的资源浪费。
    • 改进的部署灵活性:RPA v3 将编译为 3 个子内核,解锁了对纯预填充、纯解码和混合批处理的支持。这种设计通过在运行时将正确的子内核与适当的请求配对,显著节省了直接内存访问(DMA)和计算资源。
    • 这也带来了额外的好处,即解锁了更复杂的部署模式,例如解耦推理服务。
  4. 尽管 RPA v2 比第一个 TPU 原型实现了显著的吞吐量提升,但它缺乏灵活性。
    • 毫不妥协:RPA v3 没有为了灵活性而牺牲性能,事实上,在 Trillium (v6e) 上,它比 RPA v2 提升了约 10% 的吞吐量。模型现在也可以在 v5p 上运行(尽管需要额外的调优)。

我们很快会撰写关于 RPA v3 的技术深度解析,敬请关注我们的文档。

重要:要点 #4:RPA v3 既灵活又高效,是开源领域生产级 Pallas 内核开发的绝佳参考。我们非常期待 TPU 友好的 MoE 和 MLA 内核能以类似的方式尽快落地开源。

单程序多数据(SPMD)

此版本引入了单程序多数据(SPMD)作为 vLLM TPU 的默认编程模型。与之前(从 GPU 范式改编而来)的多工作节点模型不同,SPMD 是 XLA 编译器的原生模型。开发者为单个庞大的设备编写代码,XLA 编译器会自动对模型和张量进行分区,并插入通信操作以实现最佳执行。

重要:要点 #5:SPMD 实现了诸如将通信与计算重叠等高级优化。SPMD 代表了向更深、更原生的 TPU 集成迈出的战略性转变,预示着通过以 TPU 为中心、编译器优先的运行模式获得更高的性能。

总结

vLLM TPU 从 2025 年 2 月的原型性能取得了长足进步,在相同工作负载上实现了近 2-5 倍的性能提升,同时提高了模型覆盖率和易用性。

重要:要点 #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

功能

  • 前缀缓存
  • 分块预填充 (Chunked Prefill)
  • 多模态输入
  • 单程序多数据 (SPMD)
  • 结构化解码
  • 投机解码 (Speculative decoding): Ngram
  • 树外(Out-of-tree)模型支持
  • 优化后的运行时采样 (top k, top p, temperature, logit 输出)
  • 量化 (权重、激活和 KV 缓存)

TPU 友好内核

  • Ragged Paged Attention V3
  • 集合通信矩阵乘法 (Collective Communication Matmul)
  • 量化矩阵乘法、注意力机制和 KV 缓存

实验性功能

  • v5p
  • 多模态 (通过 Torchax)
  • 多 LoRA
  • 投机解码: 基于树的 Eagle 3
  • 单主机 P/D 解耦推理服务

后续计划

  • SparseCore 卸载
  • 投机解码: Eagle 3, MTP
  • TPU 友好内核
    • XL MoE
    • MLA
  • 集成
    • 强化学习 (RL)
      • 单主机和多主机
      • 同位(Colocated)和解耦部署设置
      • 通过 Pathways 实现单控制器
      • 通过前缀缓存实现多采样
      • 权重同步和重分片
      • 通过数据并行实现吞吐量优化部署
      • LoRA
      • 支持工具调用、多轮对话部署
      • 查看我们的合作伙伴项目:Tunix, MaxText, SkyRL
    • 分布式
      • 多主机动态 P/D 解耦推理服务
      • 将前缀缓存卸载到 CPU 和远程存储
      • 优化数据并行注意力负载均衡
      • 查看我们的合作伙伴项目:llm-d
  • 欢迎贡献!

立即体验!

您可以在 Google Cloud 上试用,包括 Google Kubernetes Engine (GKE)Compute EngineVertex AI。有关安装说明和开发指南,请查看以下资源:

Google Cloud 教程:GKE: 此处,Vertex AI: 此处

致谢

我们要向 vLLM 社区在这一工作中提供的持续支持表示最诚挚的感谢。特别感谢 Woosuk Kwon 带头完成了 TPU 的 V0 实现,并继续支持我们不断壮大的团队。我们还要特别感谢 Simon MoRobert ShawMichael GoinYanping Huang 在整个工作中提供的宝贵指导。还要特别感谢 Nicolo LucchesiAlexander MatveevAkshat TripathiSaheli Bhattacharjee,感谢他们作为 V1 集成和推动 Cloud Next 落地不可或缺的一份子。