使用 TPU v5e 训练模型

TPU v5e 的每个 Pod 占用空间更小(256 个芯片),经过优化,成为适用于 Transformer、文本转图片和卷积神经网络 (CNN) 训练、微调和服务的高价值产品。如需详细了解如何使用 Cloud TPU v5e 进行服务,请参阅使用 v5e 进行推断

如需详细了解 Cloud TPU v5e TPU 硬件和配置,请参阅 TPU v5e

开始使用

以下部分介绍了如何开始使用 TPU v5e。

请求配额

您需要有配额才能使用 TPU v5e 进行训练。按需 TPU、预留的 TPU 和 TPU Spot 虚拟机有不同的配额类型。如果您将 TPU v5e 用于推理,则需要单独的配额。如需详细了解配额,请参阅配额。如需申请 TPU v5e 配额,请与 Cloud 销售团队联系。

创建 Google Cloud 账号和项目

您需要拥有 Google Cloud 账号和项目才能使用 Cloud TPU。如需了解详情,请参阅设置 Cloud TPU 环境

创建 Cloud TPU

最佳实践是使用 queued-resource create 命令将 Cloud TPU v5e 预配为已排队的资源。如需了解详情,请参阅管理已排队的资源

您还可以使用 Create Node API (gcloud compute tpus tpu-vm create) 来预配 Cloud TPU v5e。如需了解详情,请参阅管理 TPU 资源

如需详细了解可用于训练的 v5e 配置,请参阅用于训练的 Cloud TPU v5e 类型

框架设置

本部分介绍了结合使用 JAX 或 PyTorch 与 TPU v5e 进行自定义模型训练的一般设置过程。

如需查看推理设置说明,请参阅 v5e 推理简介

定义一些环境变量:

export PROJECT_ID=your_project_ID
export ACCELERATOR_TYPE=v5litepod-16
export ZONE=us-west4-a
export TPU_NAME=your_tpu_name
export QUEUED_RESOURCE_ID=your_queued_resource_id

JAX 设置

如果切片形状大于 8 个芯片,则一个切片中会有多个虚拟机。在这种情况下,您需要使用 --worker=all 标志在一个步骤中对所有 TPU 虚拟机运行安装,而无需使用 SSH 单独登录每个虚拟机:

gcloud compute tpus tpu-vm ssh ${TPU_NAME}  \
   --project=${PROJECT_ID} \
   --zone=${ZONE} \
   --worker=all \
   --command='pip install -U "jax[tpu]" -f https://storage.googleapis.com/jax-releases/libtpu_releases.html'

命令标志说明

变量 说明
TPU_NAME 用户分配的 TPU 文本 ID,该 ID 是在分配已排队的资源请求时创建的。
PROJECT_ID Google Cloud 项目名称。使用现有项目或在设置 Google Cloud 项目时创建新项目
ZONE 如需了解支持的可用区,请参阅 Cloud TPU 区域和可用区文档。
worker 有权访问底层 TPU 的 TPU 虚拟机。

您可以运行以下命令来检查设备数量(此处显示的输出是使用 v5litepod-16 切片生成的)。此代码通过检查 JAX 是否看到 Cloud TPU TensorCore 并可以运行基本操作来测试是否已正确安装所有组件:

gcloud compute tpus tpu-vm ssh ${TPU_NAME} \
   --project=${PROJECT_ID} \
   --zone=${ZONE} \
   --worker=all \
   --command='python3 -c "import jax; print(jax.device_count()); print(jax.local_device_count())"'

输出将如下所示:

SSH: Attempting to connect to worker 0...
SSH: Attempting to connect to worker 1...
SSH: Attempting to connect to worker 2...
SSH: Attempting to connect to worker 3...
16
4
16
4
16
4
16
4

jax.device_count() 显示给定切片中的芯片总数。jax.local_device_count() 表示此切片中的单个虚拟机可访问的芯片数量。

# Check the number of chips in the given slice by summing the count of chips
# from all VMs through the
# jax.local_device_count() API call.
gcloud compute tpus tpu-vm ssh ${TPU_NAME} \
   --project=${PROJECT_ID} \
   --zone=${ZONE} \
   --worker=all \
   --command='python3 -c "import jax; xs=jax.numpy.ones(jax.local_device_count()); print(jax.pmap(lambda x: jax.lax.psum(x, \"i\"), axis_name=\"i\")(xs))"'

输出将如下所示:

SSH: Attempting to connect to worker 0...
SSH: Attempting to connect to worker 1...
SSH: Attempting to connect to worker 2...
SSH: Attempting to connect to worker 3...
[16. 16. 16. 16.]
[16. 16. 16. 16.]
[16. 16. 16. 16.]
[16. 16.