使用 Ray 扩缩机器学习工作负载
本文档详细介绍了如何在 TPU 上使用 Ray 和 JAX 运行机器学习 (ML) 工作负载。将 TPU 与 Ray 搭配使用有两种不同的模式:以设备为中心的模式 (PyTorch/XLA) 和以主机为中心的模式 (JAX)。
本文档假定您已设置 TPU 环境。如需了解详情,请参阅以下资源:
- Cloud TPU:设置 Cloud TPU 环境和管理 TPU 资源
- Google Kubernetes Engine (GKE):在 GKE Autopilot 中部署 TPU 工作负载或在 GKE Standard 中部署 TPU 工作负载
以设备为中心的模式 (PyTorch/XLA)
以设备为中心的模式保留了经典 PyTorch 的大部分程序化样式。在此模式下,您可以添加新的 XLA 设备类型,该类型的工作方式与任何其他 PyTorch 设备一样。每个单独的进程都与一个 XLA 设备进行交互。
如果您已经熟悉带有 GPU 的 PyTorch,并且想要使用类似的编码抽象,则此模式非常适合您。
以下部分介绍了如何在不使用 Ray 的情况下在一个或多个设备上运行 PyTorch/XLA 工作负载,以及如何使用 Ray 在多个主机上运行同一工作负载。