Cloud TPU 多切片概览
Cloud TPU 多切片是一项全栈性能扩缩技术,可让训练作业在单个切片中或在多个 Pod 中的切片上使用多个 TPU 切片,并采用标准数据并行处理。使用 TPU v4 芯片时,这意味着训练作业可以在单次运行中使用超过 4096 个芯片。对于需要的芯片数少于 4096 个的训练作业,单个切片可以提供最佳性能。不过,多个较小的切片更容易获得,因此当将多切片与较小的切片搭配使用时,启动时间会更短。

在多切片配置中部署时,每个切片中的 TPU 芯片通过芯片间互连 (ICI) 进行通信。不同切片中的 TPU 芯片通过将数据传输到 CPU(主机)来进行通信,而 CPU 又会通过数据中心网络 (DCN) 传输数据。如需详细了解如何使用多切片进行扩容,请参阅如何使用多切片将 AI 训练扩容到多达数万个 Cloud TPU 芯片。

开发者无需编写代码即可实现芯片间 DCN 通信。XLA 编译器会为您生成该代码,并将通信与计算重叠,以实现最佳性能。
概念
- 加速器类型
- 构成多切片的每个 TPU 切片的形状。多切片请求中的每个切片都采用相同的加速器类型。加速器类型由 TPU 类型(v4 或更高版本)和紧随其后的 TensorCore 数量组成。例如,
v5litepod-128指定一个具有 128 个 TensorCore 的 TPU v5e。 - 自动修复
- 当切片遇到维护事件、抢占或硬件故障时,Cloud TPU 会创建新的切片。如果资源不足以创建新的切片,则在有可用硬件之前,创建操作将不会完成。创建新的切片后,多切片环境中的所有其他切片都会重启,以便继续训练。通过正确配置的启动脚本,训练脚本无需用户干预即可自动重新启动,并从最新的检查点加载和恢复。
- 数据中心网络 (DCN)
- 延迟时间较长、吞吐量较低的网络(与 ICI 相比),用于在多切片配置中连接 TPU 切片。
- Gang 调度
- 同时预配所有 TPU 切片时,保证所有切片都成功预配或都未成功预配。
- 芯片间互连
- 用于在 TPU Pod 内连接 TPU 的高速、低延迟内部链接。
- 多切片
- 两个或更多个可通过 DCN 进行通信的 TPU 芯片切片。
- 节点
- 在多切片上下文中,节点是指单个 TPU 切片。多切片中的每个 TPU 切片都有一个节点 ID。
- 启动脚本
- 每次启动或重新启动虚拟机时运行的标准 Compute Engine 启动脚本。对于多切片,该值在 QR 创建请求中指定。如需详细了解 Cloud TPU 启动脚本,请参阅管理 TPU 资源。
- Tensor
- 一种数据结构,用于在机器学习模型中表示多维数据。
- Cloud TPU 容量的类型
您可以使用不同类型的容量创建 TPU(请参阅 TPU 定价方式中的“使用选项”部分):
预留:如需使用预留,您必须与 Google 签订预留协议。创建资源时,请使用
--reserved标志。Spot:使用 Spot 虚拟机定位抢占式配额。系统可能会抢占您的资源,以便为更高优先级作业的请求留出空间。创建资源时,请使用
--spot标志。按需:定位按需配额,无需预留且不会被抢占。TPU 请求将排入 Cloud TPU 提供的按需配额队列,但无法保证有可用的资源。默认处于选中状态,无需标志。
开始使用
-
-
In the Google Cloud console, activate Cloud Shell.
At the bottom of the Google Cloud console, a Cloud Shell session starts and displays a command-line prompt. Cloud Shell is a shell environment with the Google Cloud CLI already installed and with values already set for your current project. It can take a few seconds for the session to initialize.
ici_data_parallelismici_fsdp_parallelismici_tensor_parallelism设置环境:
$ gcloud auth login $ export QR_ID=your-queued-resource-id $ export TPU_NAME=your-tpu-name $ export PROJECT=your-project-name $ export ZONE=us-central1-a $ export NETWORK_NAME=your-network-name $ export SUBNETWORK_NAME=your-subnetwork-name $ export RUNTIME_VERSION=v2-alpha-tpuv5-lite $ export ACCELERATOR_TYPE=v5litepod-16 $ export EXAMPLE_TAG_1=your-tag-1 $ export EXAMPLE_TAG_2=your-tag-2 $ export SLICE_COUNT=4 $ export STARTUP_SCRIPT='#!/bin/bash\n'
变量说明
输入 说明 QR_ID 已排队的资源的用户分配 ID。 TPU_NAME 用户分配的 TPU 名称。 项目 Google Cloud 项目名称 ZONE 指定要在其中创建资源的可用区。 NETWORK_NAME VPC 网络的名称。 SUBNETWORK_NAME VPC 网络中子网的名称 RUNTIME_VERSION Cloud TPU 软件版本。 ACCELERATOR_TYPE v4-16 EXAMPLE_TAG_1、EXAMPLE_TAG_2 … 用于标识网络防火墙的有效来源或目标的标记。 SLICE_COUNT 切片数量。最多只能有 256 个切片。 STARTUP_SCRIPT 如果您指定了启动脚本,该脚本会在 TPU 切片预配或重启时运行。 为
gcloud创建 SSH 密钥。我们建议您将密码留空(运行以下命令后,按两次 Enter 键)。如果系统提示google_compute_engine文件已存在,请替换现有版本。$ ssh-keygen -f ~/.ssh/google_compute_engine
预配 TPU:
gcloud
$ gcloud compute tpus queued-resources \ create ${QR_ID} \ --accelerator-type=${ACCELERATOR_TYPE} \ --runtime-version=${RUNTIME_VERSION} \ --node-id=${TPU_NAME} \ --zone=${ZONE} \ [--reserved |--spot]
Google Cloud CLI 不支持所有创建 QR 选项,例如标记。如需了解详情,请参阅创建 QR。
控制台
在 Google Cloud 控制台中,前往 TPU 页面:
点击创建 TPU。
在名称字段中,输入 TPU 的名称。
在可用区框中,选择您要在其中创建 TPU 的可用区。
在 TPU 类型框中,选择加速器类型。加速器类型用于指定您要创建的 Cloud TPU 的版本和大小。如需详细了解每个 TPU 版本支持的加速器类型,请参阅 TPU 版本。
在 TPU 软件版本框中,选择软件版本。创建 Cloud TPU 虚拟机时,TPU 软件版本用于指定要安装的 TPU 运行时的版本。如需了解详情,请参阅 TPU 软件版本。
点击启用排队切换开关。
在已排队资源的名称字段中,输入已排队的资源请求的名称。
点击创建以创建已排队的资源请求。
等待已排队的资源处于
ACTIVE状态,这表示工作器节点处于READY状态。已排队的资源预配开始后,可能需要一到五分钟才能完成,具体取决于已排队资源的大小。您可以使用 gcloud CLI 或 Google Cloud 控制台来检查已排队的资源请求的状态:gcloud
$ gcloud compute tpus queued-resources \ list --filter=${QR_ID} --zone=${ZONE}
控制台
在 Google Cloud 控制台中,前往 TPU 页面:
点击已排队的资源标签页。
点击已排队的资源请求的名称。
使用 SSH 连接到 TPU 虚拟机:
$ gcloud compute tpus tpu-vm ssh ${TPU_NAME} --zone=${ZONE}
将 MaxText(包含
shardings.py)克隆到 TPU 虚拟机:$ git clone https://github.com/AI-Hypercomputer/maxtext && cd maxtext
安装 Python 3.10:
$ sudo apt-get update $ sudo apt install python3.10 $ sudo apt install python3.10-venv
创建并激活虚拟环境:
$ python3 -m venv your-venv-name $ source your-venv-name/bin/activate
在 MaxText 仓库目录中,运行设置脚本以在 TPU 切片上安装 JAX 和其他依赖项。运行设置脚本需要几分钟时间。
$ bash setup.sh
运行以下命令以在 TPU 切片上运行
shardings.py。$ python3 -m pedagogical_examples.shardings \ --ici_fsdp_parallelism 4 \ --batch_size 131072 \ --embedding_dimension 2048
您可以在日志中查看结果。TPU 应每秒达到大约 260 TFLOP 的性能,或者 FLOPS 利用率高达 90% 以上!在本例中,我们选择了 TPU 高带宽内存 (HBM) 中可容纳的大致最大批次。
您可以随意探索 ICI 之外的其他分片策略,例如,您可以尝试以下组合:
$ python3 -m pedagogical_examples.shardings \ --ici_tensor_parallelism 4 \ --batch_size 131072 \ --embedding_dimension 2048
完成后,删除已排队的资源和 TPU 切片。您应在设置切片的环境中运行这些清理步骤(先运行
exit以退出 SSH 会话)。删除操作需要两到五分钟才能完成。如果您使用的是 gcloud CLI,则可以在后台运行此命令并使用可选的--async标志。gcloud
$ gcloud compute tpus queued-resources \ delete ${QR_ID} --force (--async)
控制台
在 Google Cloud 控制台中,前往 TPU 页面:
点击已排队的资源标签页。
选中已排队的资源请求旁边的复选框。
点击 删除。
- dcn_data_parallelism
- dcn_fsdp_parallelism
- dcn_tensor_parallelism
在运行程序机器上克隆 MaxText:
$ git clone https://github.com/AI-Hypercomputer/maxtext
进入仓库目录。
$ cd maxtext
为
gcloud创建 SSH 密钥,我们建议您将密码留空(运行以下命令后,按两次 Enter 键)。如果系统提示google_compute_engine文件已存在,请选择不保留现有版本。$ ssh-keygen -f ~/.ssh/google_compute_engine
添加一个环境变量以将 TPU 切片数设置为
2。
如需使用多切片,您的 TPU 资源必须作为已排队的资源进行管理。
入门示例
本教程使用 MaxText GitHub 代码库中的代码。MaxText 是一种高性能、可任意扩缩、开源且经过充分测试的基本 LLM,采用 Python 和 Jax 编写。MaxText 旨在能够在 Cloud TPU 上高效训练。
shardings.py 中的代码旨在帮助您开始尝试使用不同的并行处理选项。例如,数据并行处理、完全分片数据并行处理 (FSDP) 和张量并行处理。代码可从单切片环境扩容到多切片环境。
ICI 并行处理
ICI 是指用于连接单个切片中的 TPU 的高速互连。ICI 分片对应于切片内的分片。shardings.py 提供三个 ICI 并行处理参数:
您为这些参数指定的值决定了每个并行处理方法的分片数。
这些输入必须受到限制,以便 ici_data_parallelism * ici_fsdp_parallelism * ici_tensor_parallelism 等于切片中的芯片数。
下表展示了 v4-8 中可用的四个芯片的 ICI 并行处理的示例用户输入:
| ici_data_parallelism | ici_fsdp_parallelism | ici_tensor_parallelism | |
| 四向 FSDP | 1 | 4 | 1 |
| 四向张量并行处理 | 1 | 1 | 4 |
| 双向 FSDP + 双向张量并行处理 | 1 | 2 | 2 |
请注意,在大多数情况下,ici_data_parallelism 应保留为 1,因为 ICI 网络足够快,几乎总是优先使用 FSDP 而不是数据并行处理。
此示例假定您熟悉如何在单个 TPU 切片上运行代码,例如使用 JAX 在 Cloud TPU 虚拟机上运行计算。此示例展示了如何在单个切片上运行 shardings.py。
使用 DCN 并行处理进行多切片分片
shardings.py 脚本接受三个用于指定 DCN 并行处理的参数,这些参数对应于每种数据并行处理类型的分片数:
这些参数的值必须受到限制,以便 dcn_data_parallelism * dcn_fsdp_parallelism * dcn_tensor_parallelism 等于切片数。
例如,对于两个切片,请使用 --dcn_data_parallelism = 2。
| dcn_data_parallelism | dcn_fsdp_parallelism | dcn_tensor_parallelism | 切片数 | |
| 双向数据并行处理 | 2 | 1 | 1 | 2 |
dcn_tensor_parallelism 应始终设置为 1,因为 DCN 不适合此类分片。对于 v4 芯片上的典型 LLM 工作负载,dcn_fsdp_parallelism 也应设置为 1,因此 dcn_data_parallelism 应设置为切片数,但这取决于应用。
随着切片数量的增加(假设您保持切片大小和每个切片的批次不变),数据并行处理量也会增加。
在多切片环境中运行 shardings.py
您可以在多切片环境中使用 multihost_runner.py 运行 shardings.py,也可以在每个 TPU 虚拟机上运行 shardings.py。在这里,我们使用 multihost_runner.py。以下步骤与 MaxText 仓库中的使用入门:对多个切片进行快速实验中的步骤非常相似,只不过这里我们运行的是 shardings.py,而不是 train.py 中更复杂的 LLM。
multihost_runner.py 工具针对快速实验进行了优化,可重复使用相同的 TPU。由于 multihost_runner.py 脚本依赖于长期有效的 SSH 连接,因此我们不建议将其用于任何长时间运行的作业。如果您想运行较长时间的作业(例如数小时或数天),我们建议您使用 multihost_job.py。
在本教程中,我们使用“运行程序”一词来表示运行 multihost_runner.py 脚本的机器。我们使用“工作器”一词来表示构成切片的 TPU 虚拟机。您可以在本地机器上或与切片位于同一项目中的任何 Compute Engine 虚拟机上运行 multihost_runner.py。不支持在工作器上运行 multihost_runner.py。
multihost_runner.py 会自动使用 SSH 连接到 TPU 工作器。
在此示例中,您将在两个 v5e-16 切片(总共 4 个虚拟机和 16 个 TPU 芯片)上运行 shardings.py。您可以修改此示例,以便在更多 TPU 上运行。