Crea una porción de TPU de varios hosts

Aprende a crear una porción de TPU de varios hosts con un grupo de instancias administrado (MIG), conéctate a la porción y ejecuta un cálculo. En esta guía de inicio rápido, se usa la opción de consumo a pedido. Ejecuta los comandos de esta guía de inicio rápido en tu terminal local o en Cloud Shell.

Antes de comenzar

  1. Accede a tu Google Cloud cuenta de. Si eres nuevo en Google Cloud, crea una cuenta para evaluar el rendimiento de nuestros productos en situaciones reales. Los clientes nuevos también obtienen $300 en créditos gratuitos para ejecutar, probar y, además, implementar cargas de trabajo.
  2. Instala Google Cloud CLI.

  3. Si usas un proveedor de identidad externo (IdP), primero debes acceder a la gcloud CLI con tu identidad federada.

  4. Para inicializar gcloud CLI, ejecuta el siguiente comando:

    gcloud init
  5. Crea o selecciona un Google Cloud proyecto.

    Roles necesarios para seleccionar o crear un proyecto

    • Seleccionar un proyecto: Para seleccionar un proyecto, no se requiere un rol de IAM específico. Puedes seleccionar cualquier proyecto en el que se te haya otorgado un rol.
    • Crear un proyecto: Para crear un proyecto, necesitas el rol de creador de proyectos (roles/resourcemanager.projectCreator), que contiene el resourcemanager.projects.create permiso. Obtén información para otorgar roles.
    • Crea un proyecto de: Google Cloud

      gcloud projects create PROJECT_ID

      Reemplaza PROJECT_ID por un nombre para el Google Cloud proyecto de que estás creando.

    • Selecciona el Google Cloud proyecto de que creaste:

      gcloud config set project PROJECT_ID

      Reemplaza PROJECT_ID por el nombre de tu Google Cloud proyecto de.

  6. Si usas un proyecto existente en esta guía, verifica que tengas los permisos necesarios para completarla. Si creaste un proyecto nuevo, ya tienes los permisos necesarios.

  7. Verifica que la facturación esté habilitada para tu Google Cloud proyecto.

  8. Habilita la API de Compute Engine con este comando:

    Roles necesarios para habilitar las APIs

    Para habilitar las APIs, necesitas el permiso serviceusage.services.enable. Si creaste el proyecto, es probable que ya tengas este permiso a través del rol de propietario (roles/owner). De lo contrario, puedes obtener este permiso a través del rol de administrador de Service Usage (roles/serviceusage.serviceUsageAdmin). Obtén información para otorgar roles.

    gcloud services enable compute.googleapis.com
  9. Instala Google Cloud CLI.

  10. Si usas un proveedor de identidad externo (IdP), primero debes acceder a la gcloud CLI con tu identidad federada.

  11. Para inicializar gcloud CLI, ejecuta el siguiente comando:

    gcloud init
  12. Crea o selecciona un Google Cloud proyecto.

    Roles necesarios para seleccionar o crear un proyecto

    • Seleccionar un proyecto: Para seleccionar un proyecto, no se requiere un rol de IAM específico. Puedes seleccionar cualquier proyecto en el que se te haya otorgado un rol.
    • Crear un proyecto: Para crear un proyecto, necesitas el rol de creador de proyectos (roles/resourcemanager.projectCreator), que contiene el resourcemanager.projects.create permiso. Obtén información para otorgar roles.
    • Crea un proyecto de: Google Cloud

      gcloud projects create PROJECT_ID

      Reemplaza PROJECT_ID por un nombre para el Google Cloud proyecto de que estás creando.

    • Selecciona el Google Cloud proyecto de que creaste:

      gcloud config set project PROJECT_ID

      Reemplaza PROJECT_ID por el nombre de tu Google Cloud proyecto de.

  13. Si usas un proyecto existente en esta guía, verifica que tengas los permisos necesarios para completarla. Si creaste un proyecto nuevo, ya tienes los permisos necesarios.

  14. Verifica que la facturación esté habilitada para tu Google Cloud proyecto.

  15. Habilita la API de Compute Engine con este comando:

    Roles necesarios para habilitar las APIs

    Para habilitar las APIs, necesitas el permiso serviceusage.services.enable. Si creaste el proyecto, es probable que ya tengas este permiso a través del rol de propietario (roles/owner). De lo contrario, puedes obtener este permiso a través del rol de administrador de Service Usage (roles/serviceusage.serviceUsageAdmin). Obtén información para otorgar roles.

    gcloud services enable compute.googleapis.com

Roles obligatorios

Si deseas obtener los permisos que necesitas para crear un MIG que forme una porción de TPU de varios hosts, conectarte a cada VM en el MIG con SSH y ejecutar comandos, pídele a tu administrador que te otorgue los siguientes roles de IAM en tu proyecto:

Para obtener más información sobre cómo otorgar roles, consulta Administra el acceso a proyectos, carpetas y organizaciones.

También puedes obtener los permisos necesarios mediante roles personalizados o cualquier otro rol predefinido.

Crea una plantilla de instancias

Para crear una plantilla de instancias para las VMs de TPU v6e, usa el gcloud compute instance-templates create comando:

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

Crear una política de cargas de trabajo

Una política de cargas de trabajo define las propiedades físicas de tus instancias de procesamiento. En las porciones de TPU, la topología del acelerador define la disposición física de los chips TPU. Se requiere especificar una topología del acelerador para las porciones de TPU de varios hosts interconectadas.

Para crear una política de cargas de trabajo para una porción de TPU de varios hosts, usa el gcloud compute resource-policies create workload-policy comando con la marca --accelerator-topology. El siguiente comando crea una política de cargas de trabajo con una topología de 2x4:

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

Crear un MIG

Ejecuta los siguientes comandos para crear un MIG que forme una porción de TPU de varios hosts.

  1. Para crear un MIG que forme una porción de TPU de varios hosts, usa el gcloud compute instance-groups managed create comando:

    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
    

    Reemplaza PROJECT_ID por el ID del Google Cloud proyecto de.

  2. Verifica que las instancias administradas se estén ejecutando con los siguientes comandos:

Instala JAX

Instala las dependencias y el framework de JAX en un entorno virtual en todas las instancias de VM de TPU del MIG. Si tus VMs de TPU tienen instalada una versión de Python anterior a la 3.11, debes instalar Python 3.11 para ejecutar la versión más reciente de JAX.

  1. Verifica qué versión de Python se ejecuta en tus VMs de TPU:

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

    Si la versión es anterior a Python 3.11, instala 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. Crea un entorno virtual:

    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. Instala JAX en el entorno virtual:

    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'
    

Ejecuta el código JAX en la porción

Para ejecutar el código JAX en una porción de TPU, debes ejecutar el código en cada host en la porción de TPU. La llamada a la función jax.device_count() deja de responder hasta que se llama en cada host de la porción. En el siguiente ejemplo, se muestra cómo ejecutar un cálculo de JAX en una porción de TPU.

Prepara el código

Crea un archivo llamado example.py en cada instancia:

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"

Ejecuta el código en la porción

Ejecuta el programa example.py en cada TPU VM de la porción:

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'

El resultado debería ser similar al siguiente ejemplo:

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

Limpia

Para evitar que se apliquen cargos a tu Google Cloud cuenta de por los recursos que usaste en esta página, borra el Google Cloud proyecto de que tiene los recursos.

Como alternativa, si deseas conservar tu proyecto, puedes borrar solo el MIG y todas las VMs del grupo con el gcloud compute instance-groups managed delete comando:

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

¿Qué sigue?