Google 深度技术文章,讲解 Ray Serve、Ray Data 在 TPU 上的自动调度和数据加载优化。对分布式 AI 基础设施工程师有参考价值。
TL;DR:这是两篇系列文章的第 2 篇。第 1 篇介绍了一个必须理解的硬件概念,以及底层的两个层次(GKE 和 Ray Core)。本篇将介绍实际用于构建应用的库:Ray Serve、Ray Data 和 Ray Train。
如果你是第一次看到本篇,先快速回顾一下。在 TPU 上运行 Ray,归根结底需要注意一点:TPU 芯片会以固定分组连接在一起,这种分组称为 slice(切片)。同一 slice 中的 host VM 通过名为 ICI 的高速链路互联。多 host 模型必须完整地落在同一个 slice 上,否则各个 worker 将无法相互通信,任务也会一直卡住。
Google Kubernetes Engine(GKE)通过 Ray Operator add-on 配置 slice 并为其中的 host 添加标签,而 Ray Core 提供的原语 slice_placement_group() 可以一次性预留整个 slice。你只需要声明 topology(即 slice 的形状,例如用 4x4 表示 16 个芯片),下面介绍的这些库便会替你处理 placement。
由 Core 在底层负责 placement 后,所有这些库都遵循同一种模式:声明 topology,让 Core 预留 slice。不同库之间唯一的区别,是你要在哪里声明它。我们将按照大多数团队采用它们的顺序展开,首先从 serving 讲起。
Serving 通常是大多数团队的起点。一个需要多块 GPU 才能容纳的模型,可能只需一个 TPU host 就能运行;而且对于 inference,TPU 往往是供应更充足、成本效益也更高的选择。Ray Serve 提供了常见的自动扩缩容、负载均衡和多模型组合能力;在 TPU 上,它通过高吞吐量引擎 vLLM 来提供 LLM serving。
抱歉,你的浏览器不支持播放此视频。
困难的情况是:模型大到一个 host 无法容纳,例如一个通过 tensor parallel 分片到 16 个芯片上的模型。这正是 Serve 只需增加一个字段 topology 就能解决的问题。
accelerator_type: TPU-V6E
accelerator_config:
kind: tpu
topology: "4x4"
这个字段值得深入理解,因为配置错误正是多 host TPU 最典型的故障来源。设置 topology 后,Serve 的 TPU backend 会跳过通常在前期创建的 placement group,转而交给 replica 处理,由 replica 在启动时创建 slice placement group。正是这种延后处理,确保了 tensor-parallel 模型的所有 worker 都位于同一个共享的 ICI mesh 上。
如果不设置这个字段,Serve 就会回退到按芯片创建 bundle。对于多 host 模型,这些 bundle 可能被分散到两个 slice 上;由于不同 slice 之间不存在 ICI,worker 永远无法完成第一次 collective。此时你不会看到 crash,而是会看到 deployment 永远停留在 DEPLOYING 状态,同时不断消耗 TPU-hours,四处排查一个实际上只是少写了一行 YAML 的问题。因此请记住:topology 字段至关重要。
在实践中,你需要基于已发布的 vLLM TPU image 部署一个 RayService(生产环境中建议使用 RayService,而不是原始的 RayCluster),等待它进入 Running 状态,然后使用 curl 请求 endpoint。GKE 官方教程涵盖了在 v5e 上运行 Llama 3 8B 和 Mistral 7B、在 v6e 上运行 Llama 3.1 70B,以及运行 Stable Diffusion。get-started 示例中的 serve 步骤完整演示了从头到尾的部署过程。
高速 accelerator 能发挥多大价值,取决于你能否持续向它输送数据。TPU 的速度非常快,普通的 loader 很容易成为瓶颈。这正是 iter_jax_batches() 要解决的问题。它提供的 batch 已经是 JAX array,并且已经完成 device sharding。因此,无论是 training input pipeline,还是大规模 batch-inference 任务,都可以直接从 Ray Data pipeline 获取数据,不会因为在 host 端执行 NumPy 到 JAX 的复制而拖慢每个 step。
ds = ray.data.read_parquet("gs://my-bucket/train/")
for batch in ds.iter_jax_batches(batch_size=1024):
# batch arrives as device-sharded JAX arrays, ready for the training step
loss = train_step(batch)
iter_jax_batches API 会替你完成 device sharding。对于最后一个不规则 batch,也就是大小不是指定 batch size 整数倍的 batch,它允许你明确选择 drop、pad 或 raise,而不是在运行三小时后才遇到 shape error。
你可以将它用作 JaxTrainer 任务的输入端;它本身也同样适合用于在一个 TPU slice 上,对大型数据集执行离线 batch inference。该功能最近刚刚加入 Ray,get-started 示例中的 data 步骤使用它完成 dataset 准备和 batch inference。
过去,在 TPU 上使用 Ray 时,training 是最令人困惑的部分,因为你既要处理 topology,还必须在代码中考虑 slice shape。JaxTrainer 解决了这个问题。它将 Ray Train 的 training loop 能力,包括 checkpointing、fault tolerance 和多 slice 横向扩展,引入 JAX。JAX 是 Google 的 array 与自动微分库,也是 TPU 的原生 framework。你只需向它提供一个 training function 和一个 slice shape,Ray 就会在每个 host 上启动一个 worker,将它们连接成一个 mesh,并在每个 worker 上运行你的函数。
from ray.train import ScalingConfig
from ray.train.v2.jax import JaxTrainer
def train_loop_per_worker(config):
import jax # import jax INSIDE the worker fn (TPU requirement)
# ... your JAX/Flax training step runs here, once per host ...
trainer = JaxTrainer(
train_loop_per_worker=train_loop_per_worker,
scaling_config=ScalingConfig(
use_tpu=True,
topology="4x4", # the slice shape, NOT a chip count
accelerator_type="TPU-V6E",
),
)
trainer.fit()
为了节省调试时间,这段代码中有两点需要牢记。第一,import jax 必须放在 train_loop_per_worker 内部,而不能放在文件顶部,因为每个 worker 都要在自己的 TPU context 中初始化 JAX。如果在 module scope 中导入它,你会在第一个 step 开始之前就遇到令人费解的 device-init error。
第二,topology="4x4" 就是完整的 placement 声明。过去需要用一整段手写的协调代码才能完成的事情,现在只需要这一行。把它和 GPU 版的 JaxTrainer 或 TorchTrainer 放在一起比较,真正的区别只有 use_tpu=True,以及用 topology 代替 GPU 数量。
剩下的部分直接运行即可。由于 Ray Train 负责管理整个 loop,你会自动获得 checkpointing 和 fault-tolerant restart。这些能力可以让运行在 preemptible capacity 上的长时间 TPU 任务真正执行完毕。当一个 slice 不够用时,topology 还可以扩展到多个 slice,由 Ray 负责不同 slice 之间的协调。get-started 示例中的 train 步骤提供了一个完整的 JaxTrainer DPO 训练任务。
作为一等 accelerator 支持的一部分,Ray 现在正式发布 rayproject/ray:*-tpu image,其中已经安装了 JAX/TPU 软件栈(jax[tpu]、flax、optax、orbax-checkpoint)和 profiling 工具,因此你无须再手动组装一套可用的 TPU 环境。只需以带有 -tpu tag 的 image 作为基础 image 即可。
在监控方面,Ray Dashboard,也就是 Ray 内置的 cluster 和任务状态 Web UI,现在会在 Cluster 标签页中,将 TPU utilization 和 memory 与 CPU、GPU 指标一起展示。ray.util.tpu.init_jax_profiler() 还会提供一个 per-worker JAX profiler,供 Dashboard 连接使用。
在这份关于 TPU 上运行 Ray 的开发者指南中,我们完整介绍了从 Ray 如何在 TPU 上运行,到如何运行 AI workload 的整个过程。
第 1 篇说明了在 TPU 上运行 Ray 归根结底只有一个关键注意事项:必须把多 host 模型放在单个完整的 slice 上;GKE 通过 Ray Operator add-on、Ray Core 通过 slice_placement_group() 替你处理这件事。
本篇则在此基础上加入了 AI library:Ray Serve 只需一个 accelerator_config.topology 字段,就能通过 gang scheduling 将多 host 模型调度到同一个 slice;Ray Data 通过 iter_jax_batches() 向 slice 提供 JAX-native batch;JaxTrainer 则通过一个 ScalingConfig 运行分布式 training loop。它还是你已经在 GPU 上使用的那个 Ray,只是现在也能运行在 TPU 上。
未来还会有更多进展。Google Cloud 上的 Ray 团队正在进一步拓展 TPU 支持:更深入的 Ray Data 与 Ray LLM TPU 集成、使用多 host TPU 的 SkyRL 来支持 reinforcement learning 和 post-training,以及动态 super/sub-slice 支持,都已经列入 roadmap。
对于你自己的下一步,我建议:克隆 get-started 示例,启动 cluster,然后运行 serve、data 或 train。或者,只需在一个 cluster 上启用 --enable-ray-operator,再在一个小型 slice 上运行一个 Ray task,亲自看看它如何工作。你不需要先成为 TPU 专家才能使用 TPU,直接试试看就好。
感谢阅读!如果你还有其他问题或反馈,欢迎通过社交媒体(LinkedIn、X)联系。
可运行示例:kubernetes-engine-samples 中的 Ray on TPU get-started,其中包含本文 serve、data 和 train 步骤对应的可运行代码(在 v6e slice 上运行 Qwen3-4B)。
在 GKE 上使用 KubeRay 和 TPU 提供 LLM serving。
开始使用 JAX 进行分布式训练。
在 Ray Dashboard 中查看 TPU metrics。
Docker Hub 上的 rayproject/ray。
第一次接触这些内容?第 1 篇介绍了 slice、GKE 和 Ray Core,它们是上述所有能力赖以构建的基础。