TPU-Slice mit mehreren Hosts erstellen

Hier erfahren Sie, wie Sie mit einer verwalteten Instanzgruppe (Managed Instance Group, MIG) einen TPU-Slice mit mehreren Hosts erstellen, eine Verbindung zum Slice herstellen und eine Berechnung ausführen. In dieser Kurzanleitung wird die Option für die On-Demand-Nutzung verwendet. Führen Sie die Befehle in dieser Kurzanleitung in Ihrem lokalen Terminal oder in Cloud Shell aus.

Hinweis

  1. Installieren Sie die Google Cloud CLI.

  2. Konfigurieren Sie die gcloud CLI für die Verwendung Ihrer föderierten Identität.

    Weitere Informationen finden Sie unter Mit Ihrer föderierten Identität in der gcloud CLI anmelden.

  3. Führen Sie den folgenden Befehl aus, um die gcloud CLI zu initialisieren:

    gcloud init
  4. Erstellen Sie ein neues Projekt oder wählen Sie ein vorhandenes Projekt Cloud de Confiance aus.

    Erforderliche Rollen zum Auswählen oder Erstellen eines Projekts

    • Projekt auswählen: Für die Auswahl eines Projekts ist keine bestimmte IAM-Rolle erforderlich. Sie können ein beliebiges Projekt auswählen, für das Ihnen eine Rolle zugewiesen wurde.
    • Projekt erstellen: Zum Erstellen eines Projekts benötigen Sie die Rolle „Projektersteller“ (roles/resourcemanager.projectCreator), die die resourcemanager.projects.create Berechtigung enthält. Rollen zuweisen.
    • Projekt erstellen: Cloud de Confiance

      gcloud projects create PROJECT_ID

      Ersetzen Sie PROJECT_ID durch einen Namen für das Cloud de Confiance Projekt, das Sie erstellen.

    • Wählen Sie das von Ihnen erstellte Cloud de Confiance Projekt aus:

      gcloud config set project PROJECT_ID

      Ersetzen Sie PROJECT_ID durch den Namen Ihres Cloud de Confiance Projekts in.

  5. Wenn Sie für diese Anleitung ein vorhandenes Projekt verwenden, prüfen Sie, ob Sie die erforderlichen Berechtigungen haben. Wenn Sie ein neues Projekt erstellt haben, haben Sie bereits die erforderlichen Berechtigungen.

  6. Prüfen Sie, ob für Ihr Cloud de Confiance Projekt die Abrechnung aktiviert ist.

  7. Aktivieren Sie die Compute Engine API:

    Erforderliche Rollen zum Aktivieren von APIs

    Zum Aktivieren von APIs benötigen Sie die Berechtigung serviceusage.services.enable. Wenn Sie das Projekt erstellt haben, haben Sie diese Berechtigung wahrscheinlich bereits über die Rolle „Inhaber“ (roles/owner). Andernfalls können Sie diese Berechtigung über die Rolle „Service Usage-Administrator“ (roles/serviceusage.serviceUsageAdmin) erhalten. Rollen zuweisen.

    gcloud services enable compute.googleapis.com

Erforderliche Rollen

Bitten Sie Ihren Administrator, Ihnen die folgenden IAM-Rollen für Ihr Projekt zuzuweisen, um die Berechtigungen zu erhalten, die Sie zum Erstellen einer MIG, die einen TPU-Slice mit mehreren Hosts bildet, zum Herstellen einer SSH-Verbindung zu jeder VM in der MIG und zum Ausführen von Befehlen benötigen:

Weitere Informationen zum Zuweisen von Rollen finden Sie unter Zugriff auf Projekte, Ordner und Organisationen verwalten.

Sie können die erforderlichen Berechtigungen auch über benutzerdefinierte Rollen oder andere vordefinierte Rollen erhalten.

Instanzvorlage erstellen

Verwenden Sie den gcloud compute instance-templates create Befehl, um eine Instanzvorlage für TPU v6e-VMs zu erstellen:

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

Arbeitslastrichtlinie erstellen

Eine Arbeitslastrichtlinie definiert die physischen Eigenschaften Ihrer Compute-Instanzen. In TPU-Slices definiert die Beschleunigertopologie die physische Anordnung der TPU-Chips. Die Angabe einer Beschleunigertopologie ist für miteinander verbundene TPU-Slices mit mehreren Hosts erforderlich.

Verwenden Sie den gcloud compute resource-policies create workload-policy Befehl mit dem --accelerator-topology Flag, um eine Arbeitslastrichtlinie für einen TPU-Slice mit mehreren Hosts zu erstellen. Der folgende Befehl erstellt eine Arbeitslastrichtlinie mit einer 2x4-Topologie:

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

Eine MIG erstellen.

Führen Sie die folgenden Befehle aus, um eine MIG zu erstellen, die einen TPU-Slice mit mehreren Hosts bildet.

  1. Verwenden Sie den gcloud compute instance-groups managed create Befehl, um eine MIG zu erstellen, die einen TPU-Slice mit mehreren Hosts bildet:

    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
    

    Ersetzen Sie PROJECT_ID durch die Cloud de Confiance by S3NS Projekt-ID Ihres Projekts.

  2. Optional können Sie mit den folgenden Befehlen prüfen, ob die verwalteten Instanzen ausgeführt werden:

JAX installieren

Installieren Sie Abhängigkeiten und das JAX-Framework in einer virtuellen Umgebung auf allen TPU-VM-Instanzen in der MIG. Wenn auf Ihren TPU-VMs eine Python-Version vor 3.11 installiert ist, müssen Sie Python 3.11 installieren, um die neueste Version von JAX auszuführen.

  1. Prüfen Sie, welche Python-Version auf Ihren TPU-VMs ausgeführt wird:

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

    Wenn die Version älter als Python 3.11 ist, installieren Sie 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. Virtuelle Umgebung erstellen:

    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. Installieren Sie JAX in der virtuellen Umgebung:

    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-Code auf dem Slice ausführen

Um JAX-Code auf einem TPU-Slice auszuführen, müssen Sie den Code auf jedem Host im TPU-Slice ausführen. Der Funktionsaufruf jax.device_count() reagiert erst, wenn er auf jedem Host im Slice aufgerufen wird. Das folgende Beispiel zeigt, wie Sie auf einem TPU-Slice eine JAX-Berechnung ausführen.

Code vorbereiten

Erstellen Sie auf jeder Instanz eine Datei mit dem Namen 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"

Den Code auf dem Slice ausführen

Führen Sie das Programm example.py auf jeder TPU-VM im Slice aus:

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'

Die Ausgabe sollte in etwa so aussehen:

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

Bereinigen

Löschen Sie das Projekt von zusammen mit den Ressourcen, damit Ihrem Cloud de Confiance Konto von die auf dieser Seite verwendeten Ressourcen nicht in Rechnung gestellt werden. Cloud de Confiance

Wenn Sie Ihr Projekt beibehalten möchten, können Sie alternativ nur die MIG und alle VMs in der Gruppe mit dem gcloud compute instance-groups managed delete Befehl löschen:

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

Nächste Schritte