在 TPU 切片上运行 JAX 代码

在运行本文档中的命令之前,请确保已按照设置账号和 Cloud TPU 项目中的说明操作。

在单个 TPU 板上运行 JAX 代码后,您可以通过在 TPU 切片上运行代码来扩容代码。TPU 切片是通过专用高速网络连接相互连接的多个 TPU 板。本文档介绍了如何在 TPU 切片上运行 JAX 代码;如需了解更深入的信息,请参阅在多主机和多进程环境中使用 JAX

所需的角色

如需获得创建 TPU 并使用 SSH 连接到该 TPU 所需的权限,请让您的管理员为您授予项目的以下 IAM 角色:

如需详细了解如何授予角色,请参阅管理对项目、文件夹和组织的访问权限

您也可以通过自定义角色或其他预定义角色来获取所需的权限。

创建 Cloud TPU 切片

  1. 创建一些环境变量:

    export PROJECT_ID=your-project-id
    export TPU_NAME=your-tpu-name
    export ZONE=europe-west4-b
    export ACCELERATOR_TYPE=v5litepod-32
    export RUNTIME_VERSION=v2-alpha-tpuv5-lite

    环境变量说明

    变量 说明
    PROJECT_ID 您的 Google Cloud 项目 ID。使用现有项目或创建新项目