マルチホスト TPU スライスを作成する

マネージド インスタンス グループ(MIG)を使用してマルチホスト TPU スライスを作成し、スライスに接続して計算を実行する方法について説明します。このクイックスタートでは、オンデマンド消費オプションを使用します。このクイックスタートのコマンドは、ローカル ターミナルまたは Cloud Shell で実行します。

始める前に

  1. Google Cloud CLI をインストールします。

  2. フェデレーション ID(連携 ID)を使用するように gcloud CLI を構成します。

    詳細については、連携 ID を使用して gcloud CLI にログインするをご覧ください。

  3. gcloud CLI を初期化するには、次のコマンドを実行します。

    gcloud init
  4. Cloud de Confiance プロジェクトを作成または選択します

    プロジェクトの選択または作成に必要なロール

    • プロジェクトを選択する: プロジェクトの選択に特定の IAM ロールは必要ありません。ロールが付与されているプロジェクトであれば、どのプロジェクトでも選択できます。
    • プロジェクトを作成する: プロジェクトを作成するには、resourcemanager.projects.create 権限を含むプロジェクト作成者ロール(roles/resourcemanager.projectCreator)が必要です。詳しくは、ロールを付与する方法をご覧ください。
    • Cloud de Confiance プロジェクトを作成します。

      gcloud projects create PROJECT_ID

      PROJECT_ID は、作成する Cloud de Confiance プロジェクトの名前に置き換えます。

    • 作成した Cloud de Confiance プロジェクトを選択します。

      gcloud config set project PROJECT_ID

      PROJECT_ID は、 Cloud de Confiance プロジェクトの名前に置き換えます。

  5. このガイドで既存のプロジェクトを使用する場合は、このガイドを完了するために必要な権限があることを確認します。新しいプロジェクトを作成した場合は、必要な権限がすでに付与されています。

  6. Cloud de Confiance プロジェクトに対して課金が有効になっていることを確認します

  7. Compute Engine API を有効にします。

    API を有効にするために必要なロール

    API を有効にするには、serviceusage.services.enable 権限が必要です。プロジェクトを作成した場合は、オーナーロール(roles/owner)を介してこの権限がすでに付与されている可能性があります。それ以外の場合は、Service Usage 管理者ロール(roles/serviceusage.serviceUsageAdmin)を介してこの権限を取得できます。ロールを付与する方法をご覧ください。

    gcloud services enable compute.googleapis.com

必要なロール

マルチホスト TPU スライスを形成する MIG の作成、MIG 内の各 VM への SSH 接続、コマンドの実行に必要な権限を取得するには、プロジェクトに対する次の IAM ロールを付与するよう管理者に依頼してください。

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

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

インスタンス テンプレートの作成

TPU v6e VM のインスタンス テンプレートを作成するには、gcloud compute instance-templates create コマンドを使用します。

gcloud compute instance-templates create quickstart-tpu-instance-template \
    --machine-type=ct6e-standard-4t \
    --maintenance-policy=TERMINATE \
    --image-family=ubuntu-accel-2204-amd64-tpu-v5e-v5p-v6e \
    --image-project=ubuntu-os-accelerator-images \
    --region=us-east5

ワークロード ポリシーの作成

ワークロード ポリシーは、コンピューティング インスタンスの物理プロパティを定義します。TPU スライスでは、アクセラレータ トポロジによって TPU チップの物理的な配置が定義されます。相互接続されたマルチホスト TPU スライスでは、アクセラレータ トポロジを指定する必要があります。

マルチホスト TPU スライスのワークロード ポリシーを作成するには、--accelerator-topology フラグを指定して gcloud compute resource-policies create workload-policy コマンドを使用します。次のコマンドは、2x4 トポロジのワークロード ポリシーを作成します。

gcloud compute resource-policies create workload-policy quickstart-tpu-workload-policy \
    --type=high-throughput \
    --accelerator-topology=2x4 \
    --region=us-east5

MIG を作成する

次のコマンドを実行して、マルチホスト TPU スライスを形成する MIG を作成します。

  1. マルチホスト TPU スライスを形成する MIG を作成するには、gcloud compute instance-groups managed create コマンドを使用します。

    gcloud compute instance-groups managed create quickstart-tpu-mig \
        --size=2 \
        --target-size-policy-mode=bulk \
        --template=quickstart-tpu-instance-template \
        --region=us-east5 \
        --target-distribution-shape=any-single-zone \
        --instance-redistribution-type=none \
        --default-action-on-vm-failure=do-nothing \
        --workload-policy=projects/PROJECT_ID/regions/us-east5/resourcePolicies/quickstart-tpu-workload-policy
    

    PROJECT_ID は、実際の Cloud de Confiance by S3NS プロジェクト ID に置き換えます。

  2. 必要に応じて、次のコマンドを使用して、マネージド インスタンスが実行されていることを確認します。

JAX をインストールする

MIG 内のすべての TPU VM インスタンスの仮想環境に、依存関係と JAX フレームワークをインストールします。TPU VM に 3.11 より前のバージョンの Python がインストールされている場合は、最新バージョンの JAX を実行するために Python 3.11 をインストールする必要があります。

  1. TPU VM で実行されている Python のバージョンを確認します。

    gcloud compute instance-groups managed list-instances quickstart-tpu-mig \
        --region=us-east5 \
        --uri \
    | xargs -I {} -P 0 gcloud compute ssh {} \
        --command='python3 --version'
    

    バージョンが Python 3.11 より前の場合は、Python 3.11 をインストールします。

    gcloud compute instance-groups managed list-instances quickstart-tpu-mig \
        --region=us-east5 \
        --uri \
    | xargs -I {} -P 0 gcloud compute ssh {} \
        --command='sudo apt update && \
        sudo apt install -y software-properties-common && \
        sudo add-apt-repository -y ppa:deadsnakes/ppa && \
        sudo apt update && \
        sudo apt install -y python3.11 python3.11-dev'
    
  2. 仮想環境を作成します。

    gcloud compute instance-groups managed list-instances quickstart-tpu-mig \
        --region=us-east5 \
        --uri \
    | xargs -I {} -P 0 gcloud compute ssh {} \
        --command='sudo apt install -y python3.11-venv && \
        python3.11 -m venv ~/jax_venv'
    
  3. 仮想環境に JAX をインストールします。

    gcloud compute instance-groups managed list-instances quickstart-tpu-mig \
        --region=us-east5 \
        --uri \
    | xargs -I {} -P 0 gcloud compute ssh {} \
        --command='source ~/jax_venv/bin/activate && \
        pip install --upgrade pip -q && \
        pip install jax[tpu] -f https://storage.googleapis.com/jax-releases/libtpu_releases.html -q'
    

スライスで JAX コードを実行する

TPU スライスで JAX コードを実行するには、TPU スライスの各ホストでコードを実行する必要があります。jax.device_count() 関数呼び出しは、スライスの各ホストで呼び出されるまで応答しなくなります。次の例は、TPU スライスで JAX 計算を実行する方法を示しています。

コードを準備する

各インスタンスに example.py というファイルを作成します。

gcloud compute instance-groups managed list-instances quickstart-tpu-mig \
    --region=us-east5 \
    --uri \
| xargs -I {} -P 0 gcloud compute ssh {} \
    --command="cat << 'EOF' > ~/example.py
import jax

# Initialize the slice
jax.distributed.initialize()

# The total number of TPU cores in the slice
device_count = jax.device_count()

# The number of TPU cores attached to this host
local_device_count = jax.local_device_count()

# The psum is performed over all mapped devices across the slice
xs = jax.numpy.ones(jax.local_device_count())
r = jax.pmap(lambda x: jax.lax.psum(x, 'i'), axis_name='i')(xs)

# Print from a single host to avoid duplicated output
if jax.process_index() == 0:
    print('global device count:', jax.device_count())
    print('local device count:', jax.local_device_count())
    print('pmap result:', r)
EOF"

スライスでコードを実行する

スライスの各 TPU VM で example.py プログラムを実行します。

gcloud compute instance-groups managed list-instances quickstart-tpu-mig \
    --region=us-east5 \
    --uri \
| xargs -I {} -P 0 gcloud compute ssh {} \
    --command='source ~/jax_venv/bin/activate && python3 ~/example.py'

出力例を以下に示します。

global device count: 8
local device count: 4
pmap result: [8. 8. 8. 8.]

クリーンアップ

このページで使用したリソースについて、 Cloud de Confiance アカウントに課金されないようにするには、リソースを含む Cloud de Confiance プロジェクトを削除します。

プロジェクトを保持する場合は、gcloud compute instance-groups managed delete コマンドを使用して、MIG とグループ内のすべての VM のみを削除できます。

gcloud compute instance-groups managed delete quickstart-tpu-mig --region=us-east5

次のステップ