使用 MaxText 在 TPU 虚拟机上运行强化学习训练

本教程提供了一份分步指南,介绍如何使用 MaxText(一种基于 JAX 的高性能训练堆栈,适用于大语言模型 [LLM]),在 Google Cloud 上的单个 v6e-8张量处理单元 [TPU] 虚拟机 [VM] 实例上运行强化学习 [RL] 训练。

目标

  • 设置 Cloud TPU 虚拟机实例。
  • 安装 MaxText 及其依赖项。
  • 将 Hugging Face 模型转换为 MaxText 格式。
  • 在 TPU 上运行 RL �����相对策略优化 (GRPO) 工作负载。
  • 将训练后的模型转换回 Hugging Face 格式以用于服务。

费用

在本文档中,您将使用 Google Cloud的以下收费组件:

如需根据您的预计使用情况来估算费用,请使用价格计算器。

新 Google Cloud 用户可能有资格申请免费试用。

完成本文档中描述的任务后,您可以通过删除所创建的资源来避免继续计费。如需了解详情,请参阅清理。

准备工作

  • 您需要拥有 Hugging Face 访问令牌才能使用本教程。您可以在 Hugging Face 上注册免费账号。拥有账号后,生成访问令牌:

    1. 在 Welcome to Hugging Face 页面上,点击您的账号头像,然后选择 Access tokens。
    2. 在访问令牌页面上,点击创建新令牌。
    3. 选择读取令牌类型,然后输入令牌的名称。
    4. 系统会显示您的访问令牌。将令牌保存在安全的位置。

  • 在 Hugging Face 网站上,接受您计划训练的模型的许可协议。本教程使用模型 llama3.1-8b-Instruct。

如需获得完成本教程所需的权限,请让您的管理员为您授予项目的以下 IAM 角色:

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

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

设置环境

运行以下脚本来设置环境变量:

export PROJECT="YOUR_PROJECT_ID"
export ZONE="YOUR_ZONE"
export RESERVATION="YOUR_RESERVATION_NAME"
export TPU_NAME="YOUR_TPU_NAME"

替换以下内容:

  • YOUR_PROJECT_ID:您的 Google Cloud 项目 ID
  • YOUR_ZONE:您要使用的可用区
  • YOUR_RESERVATION_NAME:您的容量预留
  • YOUR_TPU_NAME:Cloud TPU 虚拟机实例的名称

运行以下命令,通过 Google Cloud 进行身份验证:

gcloud auth login

创建 Cloud TPU 虚拟机

创建具有 8 个 v6e TPU 芯片的 Cloud TPU 虚拟机实例,并将其绑定到容量预留。

gcloud compute instances create "${TPU_NAME}" \
    --zone="${ZONE}" \
    --project="${PROJECT}" \
    --machine-type=ct6e-standard-8t \
    --image-project=ubuntu-os-accelerator-images \
    --image-family=ubuntu-accel-2204-amd64-tpu-v5e-v5p-v6e \
    --boot-disk-size=200GB \
    --maintenance-policy=TERMINATE \
    --instance-termination-action=DELETE \
    --provisioning-model=RESERVATION_BOUND \
    --reservation-affinity=specific \
    --reservation="${RESERVATION}"

创建虚拟机实例后,使用 SSH 连接到该实例。

gcloud compute ssh "${TPU_NAME}" --zone "${ZONE}" --project "${PROJECT}"

在 TPU 虚拟机实例中完成以下步骤。

安装 MaxText

更新 TPU 虚拟机实例中的系统软件包。

sudo apt update && sudo apt upgrade -y --fix-missing

安装 MaxText 所需的 Python 3.12 及其虚拟环境软件包。

sudo apt install -y build-essential cmake ninja-build

使用 uv 可加快 Python 软件包的安装速度。

curl -LsSf https://astral.sh/uv/install.sh | sh
source "${HOME}/.local/bin/env"

创建名为 maxtext_venv 的虚拟环境并将其激活。

uv python install 3.12
uv venv --python 3.12 --seed maxtext_venv
source maxtext_venv/bin/activate

安装 MaxText 以及它在训练后任务中所需的依赖项。

UV_TORCH_BACKEND=cpu uv pip install "maxtext[tpu-post-train]==0.2.3" --resolution=lowest

运行以下命令,安装其余必需的依赖项:

install_tpu_post_train_extra_deps

将模型转换为 MaxText 格式

如需以 MaxText 格式训练模型,您必须将其从 Hugging Face 格式转换为 MaxText 格式。

请提供以下值:

  • 您的 Hugging Face 访问令牌
  • 您要使用的模型的名称
  • 您希望以 MaxText 格式保存模型的目录
  • 加载和存储选项
export HF_TOKEN="YOUR_HF_TOKEN"
export MODEL_NAME='llama3.1-8b-Instruct'
export MODEL_CHECKPOINT_DIRECTORY=/dev/shm/$MODEL_NAME/mt-format/
export USE_PATHWAYS=0 # Set to 1 for Pathways, 0 for McJAX
export LAZY_LOAD_TENSORS=False # True to use lazy load, False to use eager load.

将 YOUR_HF_TOKEN 替换为您之前创建的 Hugging Face 访问令牌。

如需将模型从 Hugging Face 格式转换为 MaxText 格式,请运行以下脚本。此转换大约需要 5 分钟才能完成。

python3 -m maxtext.checkpoint_conversion.to_maxtext \
    model_name="${MODEL_NAME?}" \
    hf_access_token="${HF_TOKEN?}" \
    base_output_directory="${MODEL_CHECKPOINT_DIRECTORY?}" \
    scan_layers=True \
    use_multimodal=False \
    hardware=cpu \
    skip_jax_distributed_system=true \
    checkpoint_storage_use_zarr3=$((1 - USE_PATHWAYS)) \
    checkpoint_storage_use_ocdbt=$((1 - USE_PATHWAYS)) \
    --lazy_load_tensors="${LAZY_LOAD_TENSORS?}"

启动训练工作负载

转换过程完成后,您可以启动 RL 工作负载。

  1. 配置 RL 工作负载训练参数。

    # -- MaxText configuration --
    export BASE_OUTPUT_DIRECTORY=/dev/shm/$MODEL_NAME/post-train/
    RUN_NAME=$(date +%Y-%m-%d-%H-%M-%S)
    export RUN_NAME
    export CHIPS_PER_VM=8
    export NUM_BATCHES=50
    export MAXTEXT_CKPT_PATH=$MODEL_CHECKPOINT_DIRECTORY/0/items
    export TPU_ACCELERATOR_TYPE=v6e-8
    export TPU_WORKER_ID=0
    TPU_NAME=$(hostname)
    export TPU_NAME
    export TPU_SKIP_MDS_QUERY=1
    export TPU_WORKER_HOSTNAMES=localhost
    export TPU_TOPOLOGY=2x4
    export TPU_CHIPS_PER_HOST_BOUNDS=2,4,1
    export TPU_HOST_BOUNDS=1,1,1
  2. 启动训练作业。在 v6e-8 虚拟机实例上,此过程大约需要 10 分钟。

    python3 -m maxtext.trainers.post_train.rl.train_rl \
        model_name="${MODEL_NAME?}" \
        load_parameters_path="${MAXTEXT_CKPT_PATH?}" \
        run_name="${RUN_NAME?}" \
        base_output_directory="${BASE_OUTPUT_DIRECTORY?}" \
        chips_per_vm="${CHIPS_PER_VM?}" \
        num_batches="${NUM_BATCHES?}" \
        num_test_batches=10 \
        rollout_data_parallelism=1 \
        rollout_tensor_parallelism=-1

将训练后的模型转换回 Hugging Face 格式

训练工作负载完成后,将模型转换回 Hugging Face 格式。

  1. 设置导出路径和训练后的参数。

    export HF_EXPORT=/dev/shm/$MODEL_NAME/hf-trained/
    export HF_MODEL_NAME=llama3.1-8b
    export POST_TRAIN_PATH=$BASE_OUTPUT_DIRECTORY/$RUN_NAME/checkpoints/actor/$NUM_BATCHES/model_params
  2. 运行转换,将模型转换回 Hugging Face 格式。

    python3 -m maxtext.checkpoint_conversion.to_huggingface \
        model_name="${HF_MODEL_NAME?}" \
        load_parameters_path="${POST_TRAIN_PATH?}" \
        base_output_directory="${HF_EXPORT?}" \
        scan_layers=True \
        use_multimodal=False \
        weight_dtype=bfloat16

转换完成后,存储在 /dev/shm/$MODEL_NAME/hf-trained 中的调优模型即可供您使用。由于虚拟机重新启动后,您将无法再访问 /dev/shm 文件夹的内容,因此您应将调整后的模型移至持久性存储空间或将其上传到 Hugging Face Hub。

清理

为避免产生额外费用,请删除在本教程中创建的资源。

删除 TPU 虚拟机实例

删除 Cloud TPU 虚拟机实例。

gcloud compute instances delete "${TPU_NAME}" --zone="${ZONE}" --project="${PROJECT}" --quiet

后续步骤

  • 如需详细了解 Cloud TPU,请参阅 Cloud TPU 简介。
  • 如需详细了解 v6e-8 TPU 的架构和配置,请参阅 TPU v6e。