使用 Ray 扩缩机器学习工作负载

本文档详细介绍了如何在 TPU 上使用 Ray 和 JAX 运行机器学习 (ML) 工作负载。将 TPU 与 Ray 搭配使用有两种不同的模式:以设备为中心的模式 (PyTorch/XLA)以主机为中心的模式 (JAX)

本文档假定您已设置 TPU 环境。如需了解详情,请参阅以下资源:

以设备为中心的模式 (PyTorch/XLA)

以设备为中心的模式保留了经典 PyTorch 的大部分程序化样式。在此模式下,您可以添加新的 XLA 设备类型,该类型的工作方式与任何其他 PyTorch 设备一样。每个单独的进程都与一个 XLA 设备进行交互。

如果您已经熟悉带有 GPU 的 PyTorch,并且想要使用类似的编码抽象,则此模式非常适合您。

以下部分介绍了如何在不使用 Ray 的情况下在一个或多个设备上运行 PyTorch/XLA 工作负载,以及如何使用 Ray 在多个主机上运行同一工作负载。

创建 TPU