Criar uma fração de TPU com vários hosts

Saiba como criar uma fração de TPU com vários hosts usando um grupo gerenciado de instâncias (MIG), conectar-se à fração e executar um cálculo. Este guia de início rápido usa a opção de consumo sob demanda. Execute os comandos neste guia de início rápido no terminal local ou no Cloud Shell.

Antes de começar

  1. Instale a Google Cloud CLI.

  2. Configure a CLI gcloud para usar sua identidade federada.

    Para mais informações, consulte Fazer login na CLI gcloud com sua identidade federada.

  3. Para inicializar a CLI gcloud, execute o seguinte comando:

    gcloud init
  4. Crie ou selecione um Google Cloud projeto.

    Funções necessárias para selecionar ou criar um projeto

    • Selecionar um projeto: a seleção de um projeto não exige um papel específico do IAM. Você pode selecionar qualquer projeto em que tenha recebido um papel.
    • Criar um projeto: para criar um projeto, você precisa do papel de criador de projetos (roles/resourcemanager.projectCreator), que contém a resourcemanager.projects.create permissão. Saiba como conceder papéis.
    • Crie um Google Cloud projeto do:

      gcloud projects create PROJECT_ID

      Substitua PROJECT_ID por um nome para o Google Cloud projeto do que você está criando.

    • Selecione o Google Cloud projeto do que você criou:

      gcloud config set project PROJECT_ID

      Substitua PROJECT_ID pelo nome do Google Cloud projeto do.

  5. Se este guia estiver usando um projeto atual, verifique se você tem as permissões necessárias para concluir o guia. Se você criou um projeto, já tem as permissões necessárias.

  6. Verifique se o faturamento está ativado para o Google Cloud projeto.

  7. Ative a API Compute Engine:

    Funções necessárias para ativar APIs

    Para ativar as APIs, é necessário ter a permissão serviceusage.services.enable. Se você criou o projeto, provavelmente já tem essa permissão pelo papel de proprietário (roles/owner). Caso contrário, você pode receber essa permissão pelo papel de administrador de uso do serviço (roles/serviceusage.serviceUsageAdmin). Saiba como conceder papéis.

    gcloud services enable compute.googleapis.com

Funções exigidas

Para receber as permissões necessárias para criar um MIG que forma uma fração de TPU com vários hosts, conecte-se a cada VM no MIG usando SSH e execute comandos, peça ao administrador para conceder a você os seguintes papéis do IAM no projeto:

Para mais informações sobre a concessão de papéis, consulte Gerenciar o acesso a projetos, pastas e organizações.

Também é possível conseguir as permissões necessárias com papéis personalizados ou outros papéis predefinidos.

Criar um modelo de instância

Para criar um modelo de instância para VMs de TPU v6e, use o 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

Criar uma política de carga de trabalho

Uma política de carga de trabalho define as propriedades físicas das instâncias de computação. Em frações de TPU, a topologia do acelerador define a disposição física dos chips de TPU. A especificação de uma topologia de acelerador é necessária para frações de TPU interconectadas com vários hosts.

Para criar uma política de carga de trabalho para uma fração de TPU com vários hosts, use o gcloud compute resource-policies create workload-policy comando com a flag --accelerator-topology. O comando a seguir cria uma política de carga de trabalho com uma topologia 2x4:

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

Criar um MIG

Execute os comandos a seguir para criar um MIG que forma uma fração de TPU com vários hosts.

  1. Para criar um MIG que forma uma fração de TPU com vários hosts, use o 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
    

    Substitua PROJECT_ID pelo ID do Google Cloud projeto.

  2. Verifique se as instâncias gerenciadas estão em execução usando os seguintes comandos:

Instalar o JAX

Instale as dependências e a estrutura do JAX em um ambiente virtual em todas as instâncias de VM de TPU no MIG. Se as VMs de TPU tiverem uma versão do Python anterior à 3.11 instalada, será necessário instalar o Python 3.11 para executar a versão mais recente do JAX.

  1. Verifique qual versão do Python está em execução nas 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'
    

    Se a versão for anterior ao Python 3.11, instale o 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. Crie um ambiente 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. Instale o JAX no ambiente 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'
    

Executar o código JAX na fração

Para executar o código JAX em uma fração de TPU, é preciso executá-lo em cada host dessa fração. A chamada de função jax.device_count() para de responder até que seja chamada em cada host na fração. O exemplo a seguir mostra como executar um cálculo JAX em uma fração de TPU.

Preparar o código

Crie um arquivo chamado example.py em cada instância:

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"

Executar o código na fração

Execute o programa example.py em cada VM de TPU na fração:

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'

A saída será semelhante a esta:

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

Limpar

Para evitar cobranças na conta do Google Cloud pelos recursos usados nesta página, exclua o Google Cloud projeto do e os recursos.

Como alternativa, se você quiser manter o projeto, poderá excluir apenas o MIG e todas as VMs no grupo usando o gcloud compute instance-groups managed delete comando:

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

A seguir