Google Developers Blog

Run Ray on TPU, Part 2: Ray AI libraries

8.5内容质量

TL;DR · AI 摘要

在TPU上部署Ray AI库需通过topology字段确保多主机模型正确分配资源,避免因跨slice通信失败导致部署停滞。

核心要点

  • 设置topology字段可防止多主机模型跨slice部署,避免集体通信失败
  • Ray Serve通过vLLM引擎实现大模型服务,支持Llama 3.1 70B等超大模型
  • 生产环境推荐使用RayService而非原始RayCluster进行部署

结构提纲

按章节快速跳转。

  1. §TPU硬件限制

    TPU芯片必须分配在固定slice组内,跨slice通信无法实现

  2. ·Ray Core资源管理

    通过slice_placement_group()实现slice级资源预留

  3. Ray Serve配置

    topology字段决定模型部署的slice分配策略

  4. ·vLLM引擎集成

    支持高效的大模型推理,处理70B参数量模型

  5. 推荐使用RayService和预置TPU镜像进行生产部署

思维导图

用一张图看清主题之间的关系。

查看大纲文本(无障碍 / 无 JS 友好)
  • TPU上Ray AI库部署
    • 硬件限制
      • slice组通信限制
    • 资源管理
      • slice_placement_group()
      • topology配置
    • 核心库
      • Ray Serve (vLLM)
      • Ray Data
      • Ray Train

金句 / Highlights

值得收藏与分享的关键句。

#TPU#Ray#GKE#AI库#vLLM
打开原文

在 TPU 上运行 Ray,第 2 部分:Ray AI 库 - Google 开发者博客

Google Tag Manager (noscript)

结束 Google Tag Manager (noscript)

HTML

在 TPU 上运行 Ray,第 2 部分:Ray AI 库

2026 年 7 月 24 日

Ivan Nardini

AI 开发者关系

Spencer Peterson

软件工程师

分享

  • Facebook
  • Twitter
  • LinkedIn
  • 邮件

TL;DR:2 部分中的第 2 部分。第 1 部分介绍了你需要了解的硬件概念以及底层的两个层级(GKE 和 Ray Core)。本部分将展示你实际构建的库:Ray Serve、Ray Data 和 Ray Train。

回顾

如果你是第一次访问此处,快速回顾一下。在 TPU 上运行 Ray 的关键点在于:TPU 芯片被固定连接到称为 slice(通过高速链路 ICI 连接的主机虚拟机)的组中,多主机模型必须部署在完整的 slice 上,否则工作节点无法相互通信,任务将陷入停滞。

Google Kubernetes Engine(GKE)结合 Ray Operator 插件可分配 slice 并标记其主机,而 Ray Core 的原语 slice_placement_group() 可一次性保留整个 slice。你只需声明拓扑结构(如 4x4 的 slice 形状,对应 16 个芯片),底层库将自动处理部署。

由于 Core 负责底层的资源分配,所有库都遵循相同模式:声明拓扑结构,由 Core 保留 slice。不同库之间的差异仅在于你声明的维度。我们将按照团队通常采用的顺序进行讲解,首先是服务端。

在 TPU 上使用 Ray Serve

大多数团队从服务端开始。需要多个 GPU 才能运行的模型可以在单个 TPU 主机上运行,而 TPU 通常是推理任务中更易获取且成本更低的选择。Ray Serve 提供常规的自动扩展、负载均衡和多模型组合功能,在 TPU 上通过 vLLM(一个高吞吐量引擎)支持大语言模型(LLM)服务。

抱歉,你的浏览器不支持播放此视频

复杂情况出现在模型太大无法容纳单个主机(例如需要跨 16 个芯片进行张量并行的分片模型)。这时 Serve 通过一个额外字段 topology 解决问题。

code
accelerator_type: TPU-V6E
accelerator_config:
  kind: tpu
  topology: "4x4"

纯文本

已复制

这个字段值得深入理解,因为设置错误会导致经典的多主机 TPU 故障。当设置 topology 后,Serve 的 TPU 后端会跳过常规的前置 placement group,转而由副本在启动时创建 slice placement group。这种延迟机制确保张量并行模型的工作节点保留在共享的 ICI 网络中。如果省略此字段,Serve 会回退到按芯片捆绑的部署方式;对于多主机模型,这些捆绑可能分散到两个 slice 中,由于 slice 之间没有 ICI 连接,工作节点将永远无法完成首次集合操作。你不会看到崩溃,而是会看到部署卡在 DEPLOYING 状态,而你只能浪费 TPU 小时去寻找 YAML 中缺失的一行代码。记住,topology 字段就是关键差异所在。

实际部署时,你应在已发布的 vLLM TPU 镜像上部署 RayService(生产环境推荐使用 RayService 而非原始 RayCluster),等待其状态变为 Running,然后通过 curl 访问端点。官方 GKE 教程涵盖了在 v5e 上运行 Llama 3 8B 和 Mistral 7B、在 v6e 上运行 Llama 3.1 70B 以及 Stable Diffusion 的案例。入门示例的 serve 部分完整演示了从部署到运行的全过程。

在 TPU 上使用 Ray Data:通过 iter_jax_batches 为加速器提供数据

一个高速加速器的实用性取决于你能持续输入的数据量,而TPU的运行速度如此之快,以至于简单的数据加载器会成为瓶颈。这就是iter_jax_batches()方法要解决的问题。它为你提供已经转换为JAX数组且已完成设备分片的批次数据,这样训练输入流水线或大规模批量推理任务可以直接从Ray Data流水线读取数据,而不会因为主机端的NumPy到JAX复制操作而阻塞步骤。

code
ds = ray.data.read_parquet("gs://my-bucket/train/")
for batch in ds.iter_jax_batches(batch_size=1024):
    # batch以设备分片的JAX数组形式到达,随时可以用于训练步骤
    loss = train_step(batch)

Python

iter_jax_batches API会为你完成设备分片操作,并通过显式选择丢弃、填充或抛出异常的方式处理不规则的最终批次(即不是你批次大小整数倍的那一批),而不是在运行三小时后出现形状错误。

你可以将其作为JaxTrainer任务的输入端使用,它在TPU切片上独立执行大规模数据集的离线批量推理时同样非常有用。该功能最近已集成到Ray中,入门示例中的数据准备和批量推理步骤也使用了它。

TPU上的Ray Train:使用JaxTrainer进行分布式训练

在TPU上,训练曾经是Ray中令人困惑的部分,因为需要处理拓扑结构并在代码中考虑切片形状。JaxTrainer解决了这个问题。它将Ray Train的训练循环(检查点、容错、多切片扩展)引入JAX——谷歌的数组和自动微分库,也是TPU的原生框架。你只需提供一个训练函数和切片形状,Ray就会在每个主机上启动一个工作进程,将它们连接到一个统一的网格中,并在每个节点上运行你的函数。

code
from ray.train import ScalingConfig
from ray.train.v2.jax import JaxTrainer

def train_loop_per_worker(config):
    import jax            # 在worker函数内部导入jax(TPU要求)
    # ...你的JAX/Flax训练步骤在这里执行,每个主机执行一次...

trainer = JaxTrainer(
    train_loop_per_worker=train_loop_per_worker,
    scaling_config=ScalingConfig(
        use_tpu=True,
        topology="4x4",            # 切片形状,而非芯片数量
        accelerator_type="TPU-V6E",
    ),
)
trainer.fit()

在这个代码片段中,有两处需要注意以节省调试时间。import jax语句必须放在train_loop_per_worker函数内部而非文件顶部,因为每个工作进程都会在自己的TPU上下文中初始化JAX;如果在模块作用域导入,会在第一次执行前遇到难以理解的设备初始化错误。另外,topology="4x4"是完整的部署声明,这行代码取代了过去需要手动编写协调代码的代码块。与GPU的JaxTrainer或TorchTrainer并列时,唯一的实际区别是use_tpu=True和使用拓扑参数代替GPU数量。

其余部分只需正常运行即可。这是因为Ray Train掌控了整个训练循环,你将获得检查点和容错重启功能,这正是让TPU在抢占式资源上长时间运行任务最终能够完成的关键。当单个切片不足以满足需求时,拓扑结构还能扩展到多切片(Ray会处理跨切片协调)。入门示例的训练步骤就是一个完整的JaxTrainer DPO运行。

最后两个补充:TPU Docker镜像和仪表板指标

作为一流加速器支持的一部分,Ray 现在发布了官方镜像 rayproject/ray:*-tpu,其中已预装 JAX/TPU 套件(jax[tpu]、flax、optax、orbax-checkpoint)和性能分析工具,无需手动搭建 TPU 环境。你只需基于带 -tpu 后缀的镜像构建即可。

在监控方面,Ray Dashboard(Ray 内置的集群和作业状态 Web UI)现在在集群标签页中显示 TPU 利用率和内存使用情况,与 CPU 和 GPU 并列展示。通过 ray.util.tpu.init_jax_profiler() 可暴露每个工作节点的 JAX 性能分析器,Dashboard 可直接连接使用。

总结

在本篇关于 Ray on TPU 的开发者指南中,我们完整介绍了从 Ray 如何在 TPU 上运行到执行 AI 工作负载的全过程。

第一部分说明了在 TPU 上运行 Ray 的关键点:保持多主机模型在单个完整切片上运行,GKE(通过 Ray Operator 插件)和 Ray Core(通过 slice_placement_group())会自动处理这一需求。本部分则在 AI 库层面进行扩展:Ray Serve 通过单个 accelerator_config.topology 字段将多主机模型调度到单个切片;Ray Data 通过 iter_jax_batches() 向切片提供原生 JAX 批次数据;JaxTrainer 通过单个 ScalingConfig 实现分布式训练循环。你现在熟悉的 GPU 版 Ray,现已全面支持 TPU。

后续计划

更多功能正在开发中。Google Cloud 的 Ray 团队正在进一步扩展 TPU 支持:包括更深入的 Ray Data 和 Ray LLM TPU 集成、基于多主机 TPU 的 SkyRL 强化学习和训练后处理、以及动态超切片/子切片支持等均在路线图中。对于你的下一步,我建议:克隆 get-started 示例,启动集群后运行 serve、data 或 train 模块。或者直接在集群中启用 --enable-ray-operator,用小切片运行一个 Ray 任务即可体验。你无需成为 TPU 专家即可使用 TPU,只需尝试一下。感谢阅读!如有任何问题或反馈,欢迎在社交平台(LinkedIn、X)上联系。

快乐构建!

附加资源

新用户?[第一部分](#part-1)介绍了切片、GKE 和 Ray Core,这些是上述所有内容的基础。

posted in:

  • AI
  • 案例研究
  • 操作指南
  • 宣布

Previous Next

相关文章

列表

AI 云 宣布 解决方案

2026年7月16日

移动 网络 案例研究 社区

2026年7月8日

操作指南

2026年7月20日

学习

2026年6月22日

导航点