使用 MaxText 在 Qwen3-14b 模型上运行多主机监督式微调

本教程提供了一份分步指南,介绍了如何使用 MaxText 在 Cloud TPU 上对 Qwen3-14b 模型运行监督式微调 (SFT)。您将学习如何构建专用容器映像,使用加速处理套件 (XPK) 通过 Pathways 预配 Google Kubernetes Engine (GKE) 集群,以及执行多主机训练工作负载。

目标

  • 了解如何构建针对训练后优化过的自定义 MaxText 容器映像。
  • 使用 XPK 预配已启用 Pathways 的 GKE 集群。
  • 将 Qwen3 14b 模型从 Hugging Face 格式转换为 MaxText 格式。
  • 在 Cloud TPU 上运行多主机 SFT 训练工作负载。
  • 将微调后的模型转换回 Hugging Face 格式以用于提供服务。

费用

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

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

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

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

准备工作

  • 验证您的用户账号或服务账号是否具有以下角色:
    • roles/compute.admin,以创建 build 虚拟机
    • roles/artifactregistry.admin,以管理 Docker 制品库
    • roles/storage.admin,用于管理数据存储桶
    • roles/container.admin,用于创建和管理 Google Kubernetes Engine 集群
    • roles/iam.serviceAccountAdmin,用于创建工作负载服务账号
    • roles/resourcemanager.projectIamAdmin,以设置 Identity and Access Management (IAM) 政策
    • roles/iam.serviceAccountUser,以充当服务账号
  • 安装并初始化 Google Cloud CLI
  • 验证您是否已在工作站上安装 Python 3.12 或更高版本。

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

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

设置环境

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

export PROJECT="YOUR_PROJECT_ID"
export REGION="YOUR_REGION"
export ZONE="YOUR_ZONE"
export CLUSTER_NAME="YOUR_CLUSTER_NAME"
export GCS_BUCKET="YOUR_GCS_BUCKET"
export CLOUD_IMAGE_NAME="$REGION-docker.pkg.dev/$PROJECT/maxtext-images/maxtext_base:latest"
export TPU_TYPE="v6e-32"
export CLUSTER_NODEPOOL_COUNT=1
export PW_CPU_MACHINE_TYPE="c4d-standard-96"
export RESERVATION="YOUR_RESERVATION_NAME"
export MODEL_NAME="qwen3-14b"
export HF_TOKEN="YOUR_HF_TOKEN"

替换以下内容:

  • YOUR_PROJECT_ID:您的 Google Cloud 项目 ID
  • YOUR_REGION:您要使用的区域
  • YOUR_ZONE:您要使用的可用区
  • YOUR_CLUSTER_NAME:Google Kubernetes Engine 集群的名称
  • YOUR_GCS_BUCKET:Cloud Storage 存储桶的唯一名称
  • YOUR_RESERVATION_NAME:您的容量预留
  • YOUR_HF_TOKEN:您的 Hugging Face 访问令牌

准备 MaxText 容器映像

如需准备 MaxText 容器映像(包括安装所需依赖项),请完成以下步骤:

  1. 创建 Cloud Storage 存储分区,请运行以下命令:

    gcloud storage buckets create gs://$GCS_BUCKET --project=$PROJECT --location=$REGION || true
  2. 创建 Artifact Registry 代码库:

    gcloud artifacts repositories create maxtext-images \
        --repository-format=docker \
        --location=$REGION \
        --project=$PROJECT \
        --description="Docker repository for MaxText images in $REGION" || true
  3. 在代码库的根目录中创建一个名为 cloudbuild.yaml 的文件,其中包含以下内容:

    steps:
      - name: 'gcr.io/cloud-builders/docker'
        entrypoint: 'bash'
        args:
          - '-c'
          - |
            set -euo pipefail
    
            # 0. Install prerequisites (if needed)
            apt-get update && apt-get install -y curl || apk add curl || true
    
            # 1. Install uv
            curl -LsSf https://astral.sh/uv/install.sh | sh
            source $$HOME/.local/bin/env
    
            # 2. Setup Python environment and install MaxText runner
            uv venv --python 3.12 --seed maxtext_venv
            source maxtext_venv/bin/activate
            uv pip install maxtext[runner]==0.2.1 --resolution=lowest
    
            # 3. Build the Docker image (Cloud Build has Docker pre-configured)
            build_maxtext_docker_image WORKFLOW=post-training
    
            # 4. Tag the image properly
            docker tag maxtext_base_image ${_CLOUD_IMAGE_NAME}
    
    # Cloud Build automatically pushes images listed here
    images:
      - '${_CLOUD_IMAGE_NAME}'
    
    options:
      # We use a high-CPU machine to match the n4-standard-16 from the VM tutorial
      machineType: 'E2_HIGHCPU_32'
  4. 使用 Cloud Build 构建 MaxText Docker 映像:

    gcloud builds submit . \
        --project=$PROJECT \
        --region=$REGION \
        --substitutions=_CLOUD_IMAGE_NAME="${CLOUD_IMAGE_NAME}"

创建 Google Kubernetes Engine 集群

如需在 Qwen3 14b 模型上运行 SFT 训练,您需要配备 TPU 芯片的 Google Kubernetes Engine 集群。安装加速处理套件 (XPK),并创建支持 Pathways 的 GKE 集群。

# Start with creating a new virtual environment to install XPK in.
VENV_DIR=venvp3
python3 -m venv $VENV_DIR
source $VENV_DIR/bin/activate
pip install xpk==1.14.0

xpk cluster create-pathways \
  --num-slices=${CLUSTER_NODEPOOL_COUNT} \
  --tpu-type=${TPU_TYPE} \
  --pathways-gce-machine-type=${PW_CPU_MACHINE_TYPE} \
  --project=${PROJECT} \
  --zone=${ZONE} \
  --cluster=${CLUSTER_NAME} \
  --custom-cluster-arguments="--enable-ip-alias" \
  --reservation=$RESERVATION \
  --default-pool-cpu-machine-type=n4-standard-16

gcloud container clusters get-credentials $CLUSTER_NAME \
  --location=$REGION \
  --project $PROJECT

准备模型以进行训练

使用基于 CPU 的工作负载将基础模型转换为 MaxText 格式。请勿在多台机器上并行运行此任务。以下命令包含一项检查,可确保转换仅在一个 TPU 节点上运行。

xpk workload create \
  --workload "qwen-hf-to-mt" \
  --docker-image $CLOUD_IMAGE_NAME \
  --cluster ${CLUSTER_NAME} \
  --tpu-type=${TPU_TYPE} \
  --num-slices=1 \
  --project=${PROJECT} \
  --zone=${ZONE} \
  --command "[ \"\$JOB_COMPLETION_INDEX\" != \"0\" ] || \
  python3 -m maxtext.checkpoint_conversion.to_maxtext \
  model_name=${MODEL_NAME} \
  hf_access_token=${HF_TOKEN} \
  base_output_directory=gs://${GCS_BUCKET}/qwen-3-14b/max-text-format/ \
  scan_layers=True \
  use_multimodal=False \
  skip_jax_distributed_system=true \
  hardware=cpu \
  --lazy_load_tensors=True"

跟踪模型转换的进度

如需跟踪转换进度,请执行以下操作:

  1. 如需列出已在 GKE 集群上调度的 pod,请运行命令 kubectl get pod
  2. 找到名为 qwen-hf-to-mt-slice-job-0-0-HASH 的 pod。
  3. 如需实时检查 pod 的输出,请运行命令 kubectl logs -f POD_NAME

启动训练工作负载

转换过程完成后,您可以使用 XPK 启动 SFT 微调工作负载。

xpk workload create-pathways \
  --cluster=${CLUSTER_NAME} \
  --project=${PROJECT} \
  --zone=${ZONE} \
  --docker-image=$CLOUD_IMAGE_NAME \
  --workload="qwen-training" \
  --tpu-type=${TPU_TYPE} \
  --num-slices=1 \
  --command="JAX_PLATFORMS=proxy JAX_BACKEND_TARGET=grpc://127.0.0.1:29000 ENABLE_PATHWAYS_PERSISTENCE=1 \
  python3 -m maxtext.trainers.post_train.sft.train_sft \
  run_name=sft \
  base_output_directory=gs://${GCS_BUCKET}/qwen-3-14b/trained/ \
  model_name=${MODEL_NAME} \
  load_parameters_path=gs://${GCS_BUCKET}/qwen-3-14b/max-text-format/0/items/ \
  hf_access_token=${HF_TOKEN} \
  per_device_batch_size=1 \
  steps=1000 \
  profiler=xplane \
  checkpoint_storage_use_zarr3=0 \
  checkpoint_storage_use_ocdbt=0 \
  enable_single_controller=True"

监控训练工作负载

使用 XPK 命令行界面 (CLI) 监控工作负载的状态。

xpk workload list --cluster ${CLUSTER_NAME} --project ${PROJECT} --zone ${ZONE}

如需查看日志和 TPU 利用率,请使用 Google Cloud 控制台。您还可以通过运行以下命令来查看日志:

kubectl logs -f qwen-training-pathways-head-0-0-HASH

HASH 替换为 pod 名称中的数字哈希值。如需验证此哈希的值,请运行命令 kubectl get pod 并检查返回的 pod 列表。

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

训练工作负载完成后,将检查点转换回 Hugging Face 格式。

xpk workload create \
  --cluster=${CLUSTER_NAME} \
  --project=${PROJECT} \
  --zone=${ZONE} \
  --docker-image=$CLOUD_IMAGE_NAME \
  --workload="qwen-mt-to-hf" \
  --tpu-type=${TPU_TYPE} \
  --num-slices=1 \
  --command="[ \"\$JOB_COMPLETION_INDEX\" != \"0\" ] || \
  python3 -m maxtext.checkpoint_conversion.to_huggingface \
  model_name=${MODEL_NAME} \
  hf_access_token=${HF_TOKEN} \
  load_parameters_path=gs://${GCS_BUCKET}/qwen-3-14b/trained/sft/checkpoints/1000/model_params/ \
  base_output_directory=gs://$GCS_BUCKET/qwen-3-14b/hf-trained/ \
  skip_jax_distributed_system=true \
  hardware=cpu \
  scan_layers=True \
  use_multimodal=False \
  weight_dtype=bfloat16"

如需跟踪转换进度,请运行命令 kubectl logs -f qwen-mt-to-hf-slice-job-0-0-HASH,并将 HASH 替换为 pod 名称中的数字哈希值。

转换完成后,存储在 gs://$GCS_BUCKET/qwen-3-14b/hf-trained/ 中的调优模型即可使用。

清理

为避免产生额外费用,请删除在本教程中创建的资源,包括您的 Google Kubernetes Engine 集群、Cloud Storage 存储桶和 Artifact Registry 代码库。

如需删除您为本教程创建的资源,请运行以下命令:

xpk cluster delete --cluster $CLUSTER_NAME --project $PROJECT --zone $ZONE --force

gcloud storage rm --recursive gs://$GCS_BUCKET

gcloud artifacts repositories delete maxtext-images --location=$REGION --project=$PROJECT --quiet

后续步骤

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