使用 MaxText 在 TPU VM 上執行強化學習訓練

本教學課程提供逐步指南,說明如何使用 MaxText (以 JAX 為基礎的高效能訓練堆疊,適用於大型語言模型 (LLM)),在單一 v6e-8 Google Cloud Tensor 處理單元 (TPU) 虛擬機器 (VM) 執行強化學習 (RL) 訓練。

目標

  • 設定 Cloud TPU VM 例項。
  • 安裝 MaxText 及其依附元件。
  • 將 Hugging Face 模型轉換為 MaxText 格式。
  • 在 TPU 上執行 RL 群組相對政策最佳化 (GRPO) 工作負載。
  • 將訓練好的模型轉換回 Hugging Face 格式,以供使用。

費用

在本文件中,您會使用下列 Google Cloud的計費元件:

如要根據預測用量估算費用,請使用 Pricing Calculator

初次使用 Google Cloud 的使用者可能符合免費試用期資格。

完成本文所述工作後,您可以刪除建立的資源,避免繼續計費,詳情請參閱「清除所用資源」一節。

事前準備

  • 如要使用本教學課程,您需要 Hugging Face 存取權杖。您可以在 Hugging Face 申請免費帳戶。擁有帳戶後,請產生存取權杖:

    1. 在「Welcome to Hugging Face」(歡迎使用 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
  • ZONE_NAME:要使用的可用區
  • RESERVATION_NAME:運算資源預留項目
  • TPU_MACHINE_NAME:Cloud TPU VM 執行個體的名稱

執行下列指令,透過 Google Cloud 進行驗證:

gcloud auth login

建立 Cloud TPU VM

建立具有 8 個 v6e TPU 晶片的 Cloud TPU VM 執行個體,並繫結至容量預留。

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

建立 VM 執行個體後,請使用 SSH 連線至該執行個體。

gcloud compute tpus tpu-vm ssh $TPU_NAME --zone $ZONE --project $PROJECT

請在 TPU VM 執行個體中完成下列步驟。

安裝 MaxText

更新 TPU VM 執行個體中的系統套件。

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 格式,請執行下列指令碼。轉換作業大約需要五分鐘才能完成。

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 VM 執行個體上,這項作業大約需要 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 中的調整後模型即可使用。由於 VM 重新啟動時,您會失去 /dev/shm 資料夾內容的存取權,因此請將微調模型移至永久儲存空間,或上傳至 Hugging Face Hub。

清除所用資源

為避免產生額外費用,請刪除在本教學課程中建立的資源。

刪除 TPU VM 執行個體

刪除 Cloud TPU VM 執行個體。

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

後續步驟

  • 如要進一步瞭解 Cloud TPU,請參閱「Cloud TPU 簡介」。
  • 如要瞭解 v6e-8 TPU 的架構和設定詳情,請參閱 TPU v6e