使用 MaxText 在 TPU 虚拟机上运行监督式微调

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

目标

  • 设置 Cloud TPU 虚拟机实例。
  • 安装 MaxText 及其依赖项。
  • 将 Hugging Face 模型转换为 MaxText 格式。
  • 在 TPU 上运行 SFT 训练工作负载。
  • 将微调后的模型转换回 Hugging Face 格式以进行部署。

费用

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

您可使用 价格计算器 根据您的预计使用情况来估算费用。

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

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

准备工作

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

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

  • Hugging Face 网站上,接受您计划训练的模型的许可 协议。本教程使用模型 gemma3-4b

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

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

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

设置环境

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

export PROJECT="YOUR_PROJECT_ID"
export ZONE="ZONE_NAME"
export RESERVATION="RESERVATION_NAME"
export 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 $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 $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_maxtext_tpu_post_train_extra_deps
install_tpu_post_train_extra_deps

将模型转换为 MaxText 格式

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

指定您的环境变量,例如您的 Hugging Face 访问令牌、您要使用的模型的名称,以及您要以 MaxText 格式保存模型的目录。

export HF_TOKEN="YOUR_HF_TOKEN"
export MODEL_NAME='gemma3-4b'
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 格式,请运行以下脚本。此转换大约需要五分钟才能完成。

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?}

启动训练工作负载

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

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

    # -- MaxText configuration --
    export BASE_OUTPUT_DIRECTORY=/dev/shm/$MODEL_NAME/post-train/
    export RUN_NAME=$(date +%Y-%m-%d-%H-%M-%S)
    export STEPS=1000
    export PER_DEVICE_BATCH_SIZE=1
    
    # -- Dataset configuration --
    export DATASET_NAME="HuggingFaceH4/ultrachat_200k"
    export TRAIN_SPLIT="train_sft"
    export TRAIN_DATA_COLUMNS="['messages']"
    
    export MAXTEXT_CKPT_PATH=$MODEL_CHECKPOINT_DIRECTORY/0/items
  2. 启动训练作业。在 v6e-8 虚拟机实例上,此过程大约需要 10 分钟。

    python3 -m maxtext.trainers.post_train.sft.train_sft \
        run_name="${RUN_NAME?}" \
        base_output_directory="${BASE_OUTPUT_DIRECTORY?}" \
        model_name="${MODEL_NAME?}" \
        load_parameters_path="${MAXTEXT_CKPT_PATH?}" \
        per_device_batch_size="${PER_DEVICE_BATCH_SIZE?}" \
        steps="${STEPS?}" \
        hf_path="${DATASET_NAME?}" \
        train_split="${TRAIN_SPLIT?}" \
        train_data_columns="${TRAIN_DATA_COLUMNS?}" \
        profiler=xplane

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

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

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

    export HF_EXPORT=/dev/shm/$MODEL_NAME/hf-trained/
    export POST_TRAIN_PATH=$BASE_OUTPUT_DIRECTORY/$RUN_NAME/checkpoints/$STEPS/model_params
  2. 运行转换以返回到 Hugging Face 格式。

    python3 -m maxtext.checkpoint_conversion.to_huggingface \
        model_name=$MODEL_NAME \
        load_parameters_path=$POST_TRAIN_PATH \
        base_output_directory=$HF_EXPORT \
        scan_layers=True \
        use_multimodal=False \
        weight_dtype=bfloat16

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

清理

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

删除 TPU 虚拟机实例

退出 Cloud TPU 虚拟机实例,然后将其删除。

gcloud alpha compute tpus tpu-vm delete $NAME --zone=$ZONE --project=$PROJECT --quiet

后续步骤

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