MaxText を使用して TPU VM で教師ありファインチューニングを実行する

このチュートリアルでは、大規模言語モデル(LLM)用の高性能 JAX ベースのトレーニング スタックである MaxText を使用して、 Google Cloud の単一の v6e-8 Tensor Processing Unit(TPU)仮想マシン(VM)インスタンスで教師ありファインチューニング(SFT)を実行するためのステップバイステップ ガイドを提供します。

目標

  • Cloud TPU VM インスタンスを設定します。
  • MaxText とその依存関係をインストールします。
  • Hugging Face モデルを MaxText 形式に変換します。
  • TPU で SFT トレーニング ワークロードを実行します。
  • サービング用にファインチューニングされたモデルを Hugging Face 形式に戻します。

費用

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

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

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

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

始める前に

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

    1. [Welcome to Hugging Face] ページで、アカウントのアバターをクリックして [アクセス トークン] を選択します。
    2. [アクセス トークン] ページで、[新しいトークンを作成] をクリックします。
    3. [読み取り] トークン タイプを選択し、トークンの名前を入力します。
    4. アクセス トークンが表示されます。トークンは安全な場所に保存してください。

  • Hugging Face ウェブサイトで、トレーニングするモデルのライセンス契約に同意します。このチュートリアルでは、モデル gemma3-4b を使用します。

このチュートリアルを完了するために必要な権限を取得するには、プロジェクトに対する次の IAM ロールを付与するよう管理者に依頼してください。

ロールの付与については、プロジェクト、フォルダ、組織へのアクセス権の管理をご覧ください。

必要な権限は、カスタムロールや他の事前定義ロールから取得することもできます。

環境を設定する

次のスクリプトを実行して、環境変数を設定します。

export PROJECT="YOUR_PROJECT_ID"
export ZONE="YOUR_ZONE"
export RESERVATION="YOUR_RESERVATION_NAME"
export NAME="YOUR_TPU_NAME"
export NETWORK="default"

次のように置き換えます。

  • YOUR_PROJECT_ID: 実際の Google Cloud プロジェクト ID
  • YOUR_ZONE: 使用するゾーン
  • YOUR_RESERVATION_NAME: 容量予約
  • YOUR_TPU_NAME: Cloud TPU VM インスタンスの名前

次のコマンドを実行して、 Google Cloud で認証します。

gcloud auth login

Cloud TPU VM を作成する

容量予約にバインドされた 8 個の v6e TPU チップを含む Cloud TPU VM インスタンスを作成します。

gcloud compute instances create "${NAME}" \
    --zone="${ZONE}" \
    --project="${PROJECT}" \
    --network="${NETWORK}" \
    --tags="${NAME}" \
    --machine-type=ct6e-standard-8t \
    --image-project=ubuntu-os-accelerator-images \
    --image-family=ubuntu-accel-2204-amd64-tpu-v5e-v5p-v6e \
    --boot-disk-size=200GB \
    --maintenance-policy=TERMINATE \
    --instance-termination-action=DELETE \
    --provisioning-model=RESERVATION_BOUND \
    --reservation-affinity=specific \
    --reservation="${RESERVATION}"

VM インスタンスを作成したら、SSH を使用して接続します。

gcloud compute ssh "${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 build-essential cmake ninja-build

uv を使用して、Python パッケージのインストールを高速化します。

curl -LsSf https://astral.sh/uv/install.sh | sh
source "${HOME}/.local/bin/env"

maxtext_venv という名前の仮想環境を作成して有効にします。

uv python install 3.12
uv venv --python 3.12 --seed maxtext_venv
source maxtext_venv/bin/activate

MaxText と、トレーニング後のタスクに必要な依存関係をインストールします。

UV_TORCH_BACKEND=cpu 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='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 形式に変換するには、次のスクリプトを実行します。この変換には 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?}"

トレーニング ワークロードを開始する

変換プロセスが完了したら、SFT ワークロードを開始できます。

  1. SFT ワークロードのトレーニング パラメータを構成します。

    # -- MaxText configuration --
    export BASE_OUTPUT_DIRECTORY=/dev/shm/$MODEL_NAME/post-train/
    RUN_NAME=$(date +%Y-%m-%d-%H-%M-%S)
    export RUN_NAME
    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
    export TPU_ACCELERATOR_TYPE=v6e-8
    export TPU_WORKER_ID=0
    TPU_NAME=$(hostname)
    export TPU_NAME
    export TPU_SKIP_MDS_QUERY=1
    export TPU_WORKER_HOSTNAMES=localhost
    export TPU_TOPOLOGY=2x4
    export TPU_CHIPS_PER_HOST_BOUNDS=2,4,1
    export TPU_HOST_BOUNDS=1,1,1
  2. トレーニング ジョブを開始します。これには、v6e-8 VM インスタンスで約 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 に保存されているチューニング済みモデルを使用できるようになります。VM が再起動すると /dev/shm フォルダのコンテンツにアクセスできなくなるため、チューニングされたモデルを永続ストレージに移動するか、Hugging Face Hub にアップロードする必要があります。

クリーンアップ

追加料金が発生しないようにするには、このチュートリアルで作成したリソースを削除します。

TPU VM インスタンスを削除する

Cloud TPU VM インスタンスを終了して削除します。

# 1. Delete TPU instance
echo "Deleting TPU instance: ${NAME}..."
gcloud compute instances delete "${NAME}" \
    --zone="${ZONE}" \
    --project="${PROJECT}" \
    --quiet || true

# 2. Delete IAP firewall rule
echo "Deleting firewall rule: ${FIREWALL_RULE_NAME}..."
gcloud compute firewall-rules delete "${FIREWALL_RULE_NAME}" \
    --project="${PROJECT}" \
    --quiet || true

次のステップ

  • Cloud TPU の詳細については、Cloud TPU の概要をご覧ください。
  • v6e-8 TPU のアーキテクチャと構成の詳細については、TPU v6e をご覧ください。