在 TPU 上运行 Ray,第 1 部分:基础原理
2026 年 7 月 20 日 Ivan Nardini AI 开发者关系 Spencer Peterson 软件工程师

太长不看版:如果你已经在 GPU 上使用 Ray 扩展 Python,那么现在,你的代码可以通过已全面支持的官方 API 在 TPU(Tensor Processing Unit,张量处理单元,Google 的 AI 加速器芯片)上运行。你已熟悉的任务与 actor 模型、一个 JaxTrainer,以及相同的 Ray Serve 部署,都可以直接指向由 Google Kubernetes Engine(GKE)编排的 TPU。
从 Ray 2.55 版本开始,Google Cloud TPU 已成为 Ray 的一等加速器。这意味着 TPU 现已正式进入 Ray 的发布流水线,拥有官方预构建镜像以及跨核心库的支持,而不再是过去那种需要自行构建容器并依赖社区帮助的“实验性”路径。在这个“在 TPU 上运行 Ray”系列中,你将了解到一个 TPU 切片如何成为 Ray 调度的另一种加速器(第 1 部分),然后我们将逐一介绍每个库的用法(第 2 部分)。
Ray 与 TPU:快速入门
Ray 是一个分布式计算框架:你编写 Python 代码,Ray 将其以任务(无状态函数)和actor(有状态工作单元)的形式在集群上运行。对 Ray 而言,TPU 就像 GPU 一样,只是另一种可调度的资源。你请求它,Ray 就把你的工作放到上面去运行。
但有一点需要记住,然后我们继续往下讲。
TPU 芯片通过固定线路连接在一起,形成一个称为 切片(slice) 的固定组合:若干台主机(虚拟机)的芯片共享一条名为 ICI(Inter-Chip Interconnect,芯片间互连)的专用高速链路。一个多主机模型必须落入一整个完整的切片 中,否则其各工作单元无法相互通信,作业便会陷于停滞。*
如果你从 GPU 的角度来理解,可以把一个切片想象成一个多 GPU 机箱,其中高速互连(NVLink)仅存在于这个机箱内部。将你的工作单元分散到两个没有线缆相连的机箱中,那么负责同步梯度的集体操作,即全归约(all-reduce)步骤,就永远无法完成。训练就会陷于停滞。一个 TPU 切片的行为方式与之相同:ICI 就是那条线缆,它只连接到同一个切片内的芯片。

这就是“在 TPU 上运行 Ray”需要特殊处理的全部原因。必须有某种机制来保证你的所有工作单元都落在一个完整的切片上。 在使用 GPU 时你几乎无需考虑这一点;但在使用 TPU 时这至关重要,而 Ray 和 GKE(Google Kubernetes Engine,Google 的托管 Kubernetes)会为你处理这个问题。
另一个你会经常看到的词是拓扑结构(topology):它指的是一个切片的形状,例如一个 16 芯片的切片,其拓扑结构被写作 4x4。你请求的是一个拓扑结构,而不是芯片数量。
一旦你理解了 TPU 的切片和拓扑结构,现有的 Ray 技术栈和你的开发流程就保持不变,它们将在 GKE 供应的 TPU 切片上运行。下面的图表是整个系统在一张图中的示意:左侧是你编写的代码(即你已经在使用的 Ray 库);中间是 Ray Core 层,负责预留整个切片;右侧是 GKE 托管层,负责供应硬件并为其打上标签,以便 Ray 能够识别切片边界。

GKE 供应切片并为其主机打上标签,Ray Core 读取这些标签来一次性预留整个切片,而你的库调用位于其上层,只需声明一个拓扑结构,无需其他任何操作。任何地方都不需要手动编写放置代码。本部分的其余内容将介绍底部的两层,即 GKE 和 Ray Core,而第 2 部分将涵盖 Ray AI 库。
GKE 如何编排 TPU 上的 Ray
你可以通过 GKE,使用 Ray Operator 附加组件,在 TPU 上运行 Ray。
# Autopilot(全托管节点)
gcloud container clusters create-auto CLUSTER \
--enable-ray-operator --location=LOCATION
# 或 Standard(你自行管理节点池)
gcloud container clusters create CLUSTER \
--addons=RayOperator --location=LOCATION &&
gcloud container node-pools create v6e-16-slice \
--cluster=CLUSTER \
--location=LOCATION \
--machine-type=ct6e-standard-4t \
--tpu-topology=4x4 \
--num-nodes=4
上述这一条命令安装了两个对 TPU 至关重要的组件。第一个组件是 KubeRay,这是一个 Kubernetes Operator,它负责将 RayCluster、RayService 和 RayJob 的 YAML 文件转换为运行中的 Ray 集群;这和你搭配 GPU 使用的 KubeRay 是同一个。第二个是 TPU 专用的部分:Ray TPU Webhook,它会为每个 TPU 主机打上诸如 ray.io/tpu-slice-name 这样的标签,以便 Ray 能够识别哪些机器连接在同一个切片内。这个标签正是整个系统运转所依赖的那条线索。
接下来,你在清单文件中请求 TPU 的方式,与你请求任何其他节点的方式相同,通过一个 nodeSelector 来指定代次和拓扑结构,并将芯片数量作为资源进行请求。对于多主机切片,则需要额外添加一个字段:numOfHosts。
# 位于 RayCluster 的 workerGroupSpec 内部
nodeSelector:
cloud.google.com/gke-tpu-accelerator: tpu-v6e-slice # TPU 代次
cloud.google.com/gke-tpu-topology: "4x4" # 切片形状
# ... 并通过 google.com/tpu 资源限制来请求芯片
numOfHosts: 4 # 多主机:构成此切片的主机(虚拟机)数量
GKE 供应切片,Webhook 为其打上标签,Ray 读取这些标签。而你只需编写 Python 代码。一旦该附加组件启动,你就能看到正在运行的 KubeRay Operator Pod,而应用上述清单文件则会启动一个头节点 Pod,并为切片中的每个主机启动一个工作节点 Pod。入门示例中的集群步骤可通过 Terraform 供应所有这些资源。
真正让你的各工作单元保持在同一位置的是位于该层之上的一个 Ray Core 原语,即切片放置组(slice placement group),本指南的其余部分将从这里开始。
TPU 上的 Ray Core
Ray Core 是基础层,是任务与 actor 的引擎和调度器,所有其他组件都构建在其上。其 TPU 支持存在于公共的 ray.util.tpu API 中,而你需要了解的核心函数只有一个:slice_placement_group()。它将前面提到的“将我的工作单元保留在一个完整切片上”的保证,转化为一个单一调用,通过匹配 Webhook 标签,原子性地(要么获得所有主机,要么全无)预留一个整个切片。
from ray.util.tpu import slice_placement_group
from ray.util.scheduling_strategies import PlacementGroupSchedulingStrategy
# 原子性地预留一个完整的 v6e 4x4 切片(16 个芯片,分布在 4 台主机上)
spg = slice_placement_group(topology="4x4", accelerator_version="v6e")
ray.get(spg.placement_group.ready(), timeout=600)
@ray.remote(resources={"TPU": 4})
def worker(rank, world): ...
tasks = [
worker.options(
scheduling_strategy=PlacementGroupSchedulingStrategy(
placement_group=spg.placement_group)
).remote(rank=i, world=spg.num_hosts)
for i in range(spg.num_hosts)
]
需要强调的是,你很少需要亲自调用 slice_placement_group。Ray 的 AI 库(Data、Train、Serve)会为你调用它,所以在实践中,你声明一个拓扑结构,它们就会处理好切片。只有当你编写不属于 Train、Serve 或 Data 的自定义分布式工作负载时,你才需要直接使用 slice_placement_group()。有一点需要注意:此 API 是公开的,但被标记为 alpha (@PublicAPI(stability="alpha")),所以它现在立即可用,但其接口在不同版本间仍可能发生变化。
基础已奠定。接下来是各库的用法。
现在你已经拥有了完整的思维模型:一个切片必须保持完整,GKE 供应并标记它,Ray Core 将其作为一个单元来预留,这样你就永远不需要手动编写放置代码。你所构建的一切都建立在这个基础之上并复用它。
在第 2 部分中,我们将探讨如何在 TPU 上使用 Ray AI 库,包括使用 vLLM 提供 LLM 服务,使用 Ray Data 供给切片,以及使用 JaxTrainer 进行训练。