在 TPU v6e 上运行 Gemma 4 26B 的多主机强化学习训练

本教程向您展示如何使用 MaxText 和 Cluster Toolkit 在张量处理单元 (TPU) v6e-64 集群上运行多主机强化学习 (RL) 训练。您可以利用 Cluster Toolkit 执行多主机训练工作负载,并将结果导出为 Hugging Face 格式以供服务。

目标

  • 安装 Cluster Toolkit 及其依赖项。
  • 安装 MaxText 及其依赖项。
  • 部署 Cluster Toolkit 集群。
  • 将 Hugging Face 模型转换为 MaxText 格式。
  • 在 TPU v6e 集群上运行 RL 训练工作负载。
  • 将微调后的模型转换回 Hugging Face 格式以用于提供服务。

费用

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

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

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

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

准备工作

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

  1. Welcome to Hugging Face 页面上,点击您的账号头像,然后选择 Access tokens
  2. 访问令牌页面上,点击创建新令牌
  3. 选择读取令牌类型,然后输入令牌的名称。
  4. 您的访问令牌已显示。将令牌保存在安全的位置。
  • Hugging Face 网站上,接受您计划训练的模型的许可协议。本教程使用模型 gemma4-26b

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

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

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

设置环境变量

设置环境变量:

export PROJECT="YOUR_PROJECT_ID"
export REGION="YOUR_REGION"
export ZONE="YOUR_ZONE"
export CLUSTER_NAME="YOUR_CLUSTER_NAME"
export REPOSITORY_NAME="YOUR_REPOSITORY_NAME"
export GCS_BUCKET="YOUR_BUCKET_NAME"
export CLOUD_IMAGE_NAME="${REGION}-docker.pkg.dev/${PROJECT}/${REPOSITORY_NAME}/maxtext_base:latest"
export COMPUTE_TYPE="ct6e-standard-4t"
export TPU_TYPE="v6e-64"
export TOPOLOGY="8x8"
export CLUSTER_NODEPOOL_COUNT=1
export PW_CPU_MACHINE_TYPE="c4d-standard-96"
export RESERVATION="YOUR_RESERVATION_NAME"
export MODEL_NAME="gemma4-26b"
export HF_TOKEN="YOUR_HF_TOKEN"

替换以下内容:

  • YOUR_PROJECT_ID:您的 Google Cloud 项目的 ID。
  • YOUR_REGION:您要在其中部署集群的区域。
  • YOUR_ZONE:您要部署集群的可用区。
  • YOUR_CLUSTER_NAME: 您的 Google Kubernetes Engine 集群的名称。
  • YOUR_REPOSITORY_NAME:MaxText 映像的 Artifact Registry 代码库的名称。
  • YOUR_BUCKET_NAME:Cloud Storage 存储桶的全球唯一名称。
  • YOUR_RESERVATION_NAME:预留的名称。
  • YOUR_HF_TOKEN:您的 Hugging Face 访问令牌。

安装 Cluster Toolkit 依赖项

如需从 Linux 或 macOS 客户端或工作站完成本教程,请按照 Cluster Toolkit 文档中的安装依赖项中的相关步骤操作。

如果您使用的是 Cloud Shell,则可以跳过此部分。

安装 Cluster Toolkit

按照 安装 Cluster Toolkit 中的说明安装 Cluster Toolkit 的预构建包。

准备 MaxText 容器映像

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

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

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

    gcloud artifacts repositories create ${REPOSITORY_NAME} \
        --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.4 --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:
      machineType: 'E2_HIGHCPU_32'
  4. 使用 Cloud Build 构建 MaxText Docker 映像:

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

创建您的集群工具包集群

要创建和部署包含 64 个 v6e TPU 芯片的集群工具包集群,请完成以下步骤:

  1. 创建名为 gke.gcsfuse.profileUser 的自定义 Identity and Access Management (IAM) 角色:

    # The GKE TPU v6e blueprint uses GCS Fuse CSI Storage Profiles which requires a custom IAM role.
    # If this role is not already created in your project, you must create it before deploying.
    gcloud iam roles create gke.gcsfuse.profileUser \
      --project=${PROJECT} \
      --title="GKE GCSFuse Profile User" \
      --description="Allows scanning GCS buckets for objects, retrieving bucket metadata, and creating Anywhere Caches." \
      --permissions="storage.objects.list,storage.buckets.get,storage.anywhereCaches.create,storage.anywhereCaches.get,storage.anywhereCaches.list,storage.anywhereCaches.update"
  2. 创建 Cloud Storage 存储分区,请运行以下命令:

    gcloud storage buckets create gs://$GCS_BUCKET --project=$PROJECT --location=$REGION || true
  3. 默认情况下,您的集群节点池服务账号没有写入 Cloud Storage 存储桶所需的权限。如需允许节点池服务账号写入您的 Cloud Storage 存储桶,您必须为其授予 Storage Admin 角色。要授予此角色,请编辑文件 gke-tpu-v6e-advanced.yaml,更新名为 node_pool_service_accountservice-account 模块

    - id: node_pool_service_account
      source: modules/project/service-account
      settings:
        name: gke-np-sa
        project_roles:
        - logging.logWriter
        - monitoring.metricWriter
        - monitoring.viewer
        - stackdriver.resourceMetadata.writer
        - storage.admin            # Change from storage.objectViewer
        - artifactregistry.reader
  4. gke-tpu-v6e-cluster 块应用自定义设置,以覆盖默认的 IPv6 和机器类型设置:

    - id: gke-tpu-v6e-cluster
      source: modules/scheduler/gke-cluster
      use: [gke-tpu-v6e-net-0, workload_service_account]
      settings:
        enable_private_ipv6_google_access: false
        system_node_pool_disk_size_gb: $(vars.system_node_pool_disk_size_gb)
        system_node_pool_taints: []
        enable_private_endpoint: false # Allows access from authorized public IPs
        enable_pathways_for_tpus: $(vars.enable_pathways_for_tpus)
        enable_dataplane_v2: true
        configure_workload_identity_sa: true
  5. 使用蓝图 gke-tpu-v6e-advanced.yaml 部署 Cluster Toolkit 集群,并通过 --vars 标志传递所需变量:

    ./gcluster deploy examples/gke-tpu-v6e/gke-tpu-v6e-advanced.yaml \
        --vars "project_id=${PROJECT},deployment_name=${CLUSTER_NAME},region=${REGION},zone=${ZONE},num_slices=${CLUSTER_NODEPOOL_COUNT},tpu_topology=${TOPOLOGY},authorized_cidr=0.0.0.0/0,reservation=${RESERVATION:-}" \
        -l IGNORE --auto-approve -w

将模型转换为 MaxText 格式

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

  1. 创建 Cluster Toolkit 集群后,请配置 Docker:

    # Configure docker for pulling images
    gcloud auth configure-docker gcr.io --quiet
    gcloud auth configure-docker ${REGION}-docker.pkg.dev --quiet
  2. 为简化后续命令,请配置默认项目、集群和位置:

    # Configure gcluster Defaults
    ./gcluster job config set project ${PROJECT}
    ./gcluster job config set cluster ${CLUSTER_NAME}
    ./gcluster job config set location ${REGION}
  3. 将模型从 Hugging Face 格式转换为 MaxText 格式,并将其存储在您的 Cloud Storage 存储桶中:

    ./gcluster job submit \
      --name="gemma4-hf-to-mt" \
      --cluster="${CLUSTER_NAME}" \
      --project="${PROJECT}" \
      --location="${REGION}" \
      --num-slices=1 \
      --image="${CLOUD_IMAGE_NAME}" \
      --compute-type="${COMPUTE_TYPE}" \
      --topology="${TOPOLOGY}" \
      --await-job-completion \
      --command="[ \"\$JOB_COMPLETION_INDEX\" != \"0\" ] || \
        python3 -m maxtext.checkpoint_conversion.to_maxtext \
        model_name=${MODEL_NAME} \
        hf_access_token=${HF_TOKEN} \
        --hf_model_path='google/gemma-4-26b-a4b-it' \
        base_output_directory=gs://${GCS_BUCKET}/${MODEL_NAME}/max-text-format/ \
        scan_layers=True \
        use_multimodal=False \
        skip_jax_distributed_system=true \
        checkpoint_storage_use_zarr3=0 \
        checkpoint_storage_use_ocdbt=0 \
        hardware=cpu \
        --lazy_load_tensors=True"
  4. 检查转换作业的状态:

    # Use the list command to check status
    ./gcluster job list \
        --cluster ${CLUSTER_NAME} \
        --project ${PROJECT} \
        --location ${REGION}
    
    # Check progress of the job (--main-only targets the coordinator pod (Job Index 0, Pod Index 0) to avoid duplicate logs from other workers)
    ./gcluster job logs gemma4-hf-to-mt --main-only -f \
        --cluster ${CLUSTER_NAME} \
        --project ${PROJECT} \
        --location ${REGION}

启动训练工作负载

转换过程完成后,开始强化学习训练工作负载:

./gcluster job submit \
  --name="gemma4-training" \
  --cluster="${CLUSTER_NAME}" \
  --project="${PROJECT}" \
  --location="${REGION}" \
  --num-slices=1 \
  --image="${CLOUD_IMAGE_NAME}" \
  --compute-type="${COMPUTE_TYPE}" \
  --topology="${TOPOLOGY}" \
  --pathways \
  --pathways-gcs-location="gs://${GCS_BUCKET}/pathways/" \
  --env="GRPC_DNS_RESOLVER=native" \
  --pathways-proxy-env="GRPC_DNS_RESOLVER=native" \
  --pathways-server-env="GRPC_DNS_RESOLVER=native" \
  --pathways-worker-env="GRPC_DNS_RESOLVER=native" \
  --command="export VLLM_HOST_IP=\$(hostname -I | awk '{print \$1}'); \
      python3 -c \"import pathlib, tpu_inference.layers.common.fused_moe_gmm as f; p = pathlib.Path(f.__file__); p.write_text(p.read_text().replace('onehot_moe_permute_threshold: int = 0,', 'onehot_moe_permute_threshold: int = 100000,'))\"; \
      JAX_PLATFORMS=proxy,cpu ENABLE_PATHWAYS_PERSISTENCE=1 \
      python3 -m maxtext.trainers.post_train.rl.train_rl \
      run_name=rl \
      base_output_directory=gs://${GCS_BUCKET}/${MODEL_NAME}/trained/ \
      model_name=${MODEL_NAME} \
      scan_layers=False \
      load_parameters_path=gs://${GCS_BUCKET}/${MODEL_NAME}/max-text-format/0/items/ \
      hf_access_token=${HF_TOKEN} \
      num_batches=50 \
      batch_size=8 \
      rollout_tensor_parallelism=2 \
      rollout_expert_parallelism=4 \
      trainer_devices_fraction=0.5 \
      sampler_devices_fraction=0.5 \
      tokenizer_path='google/gemma-4-26b-a4b-it' \
      ici_tensor_parallelism=2 \
      ici_expert_parallelism=4 \
      hbm_utilization_vllm=0.55 \
      remat_policy=full \
      async_scheduling=False \
      allow_split_physical_axes=true \
      ragged_gather_reduce_fallback=True \
      vllm_hf_overrides='{architectures: [\"MaxTextForCausalLM\"]}' \
      vllm_additional_config=\"{'maxtext_config': {'model_name': '${MODEL_NAME}', 'allow_split_physical_axes': 'true', 'use_ragged_sort': 'false', 'ragged_gather_reduce_fallback': 'true', 'prefuse_moe_weights': 'true', 'weight_dtype': 'bfloat16'}}\""

查看培训工作的状态:

# Use the list command to check status
./gcluster job list \
    --cluster ${CLUSTER_NAME} \
    --project ${PROJECT} \
    --location ${REGION}

# Check progress of the job (--main-only targets the coordinator pod (Job Index 0, Pod Index 0) to avoid duplicate logs from other workers)
./gcluster job logs gemma4-training --main-only -f \
    --cluster ${CLUSTER_NAME} \
    --project ${PROJECT} \
    --location ${REGION}

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

训练完成后,将模型转换回 Hugging Face 格式:

./gcluster job submit \
  --name="gemma4-mt-to-hf" \
  --cluster="${CLUSTER_NAME}" \
  --project="${PROJECT}" \
  --location="${REGION}" \
  --num-slices=1 \
  --image="${CLOUD_IMAGE_NAME}" \
  --compute-type="${COMPUTE_TYPE}" \
  --topology="${TOPOLOGY}" \
  --await-job-completion \
  --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}/${MODEL_NAME}/trained/rl/checkpoints/actor/50/model_params/ \
      base_output_directory=gs://${GCS_BUCKET}/${MODEL_NAME}/hf-trained/ \
      skip_jax_distributed_system=true \
      hardware=cpu \
      scan_layers=True \
      use_multimodal=False \
      weight_dtype=bfloat16 \
      --override_model_architecture"

检查转换作业的状态:

# Use the list command to check status
./gcluster job list \
    --cluster ${CLUSTER_NAME} \
    --project ${PROJECT} \
    --location ${REGION}

# Check progress of the job (--main-only targets the coordinator pod (Job Index 0, Pod Index 0) to avoid duplicate logs from other workers)
./gcluster job logs gemma4-mt-to-hf --main-only -f \
    --cluster ${CLUSTER_NAME} \
    --project ${PROJECT} \
    --location ${REGION}

# The trained model is now available in gs://${GCS_BUCKET}/${MODEL_NAME}/hf-trained/

清理

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

./gcluster destroy ${CLUSTER_NAME} --robust
gcloud storage rm -r gs://${GCS_BUCKET}
gcloud artifacts repositories delete ${REPOSITORY_NAME} --location=${REGION} --project=${PROJECT} --quiet

# To delete the local deployment folder
rm -rf .ghpc ${CLUSTER_NAME}

后续步骤