MaxText を使用して TPU VM で強化学習トレーニングを実行する

このチュートリアルでは、MaxText を使用して、 上の単一の v6e-8 Tensor Processing Unit(TPU)仮想マシン(VM) インスタンスで 強化学習(RL)トレーニングを実行する手順について説明します。MaxText は、大規模言語モデル(LLM)用の高性能な JAX ベースのトレーニング スタックです。 Google Cloud

目標

  • Cloud TPU VM インスタンスを設定する。
  • MaxText とその依存関係をインストールする。
  • Hugging Face モデルを MaxText 形式に変換する。
  • TPU で RL Group Relative Policy Optimization(GRPO)ワークロードを実行する。
  • トレーニング済みのモデルをサービング用に Hugging Face 形式に戻す。

費用

このドキュメントでは、課金対象である次のコンポーネントを使用します。 Google Cloud

料金計算ツールを使うと、予想使用量に基づいて費用の見積もりを生成できます。

新規の Google Cloud ユーザーの方は、無料トライアルをご利用いただける場合があります。

このドキュメントに記載されているタスクの完了後、作成したリソースを削除すると、それ以上の請求は発生しません。詳細については、 クリーンアップをご覧ください。

始める前に

  • このチュートリアルを使用するには、Hugging Face アクセス トークンが必要です。Hugging Faceで無料アカウントに登録できます 。アカウントを作成したら、アクセス トークンを生成します。

    1. [Welcome to Hugging Face] ページで、 アカウントのアバターをクリックし、[**Access tokens**] を選択します。
    2. [Access tokens] ページで、[Create new token] をクリックします。
    3. [Read] トークンタイプを選択し、トークンの名前を入力します。
    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 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 形式に変換するには、次のスクリプトを実行します。この変換には約 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 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 をご覧ください。