使用 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="ZONE_NAME"
export RESERVATION="RESERVATION_NAME"
export TPU_NAME="TPU_MACHINE_NAME"

替换以下内容:

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

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

gcloud auth login

创建 Cloud TPU 虚拟机

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

gcloud alpha compute tpus tpu-vm create $TPU_NAME \
    --zone=$ZONE \
    --project=$PROJECT \
    --accelerator-type=v6e-8 \
    --version=v2-alpha-tpuv6e \
    --provisioning-model=reservation-bound \
    --reservation=$RESERVATION

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

gcloud compute tpus tpu-vm 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 python3.12 python3.12-venv

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

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

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

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

安装 MaxText 及其后训练任务所需的依赖项。

uv pip install maxtext[tpu-post-train]==0.2.2 --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/
    export RUN_NAME=$(date +%Y-%m-%d-%H-%M-%S)
    export CHIPS_PER_VM=8
    export NUM_BATCHES=50
    export MAXTEXT_CKPT_PATH=$MODEL_CHECKPOINT_DIRECTORY/0/items
  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 alpha compute tpus tpu-vm delete $TPU_NAME --zone=$ZONE --project=$PROJECT --quiet

后续步骤

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