Prepare um modelo com a TPU v5e

Com uma área de 256 chips por Pod, a TPU v5e está otimizada para ser um produto de elevado valor para a preparação, o ajuste fino e a publicação de transformadores, texto para imagem e redes neurais convolucionais (CNN). Para mais informações sobre a utilização do Cloud TPU v5e para publicação, consulte o artigo Inferência com o v5e.

Para mais informações sobre o hardware e as configurações da TPU v5e do Cloud TPU, consulte o artigo TPU v5e.

Começar

As secções seguintes descrevem como começar a usar a TPU v5e.

Quota de pedidos

Precisa de quota para usar a TPU v5e para preparação. Existem diferentes tipos de quotas para TPUs a pedido, TPUs reservadas e VMs Spot de TPU. Existem quotas separadas necessárias se estiver a usar a TPU v5e para inferência. Para mais informações sobre quotas, consulte o artigo Quotas. Para pedir quota de TPUs v5e, contacte as vendas do Google Cloud.

Crie uma Google Cloud conta e um projeto

Precisa de uma Google Cloud conta e um projeto para usar o Cloud TPU. Para mais informações, consulte o artigo Configure um ambiente de TPU na nuvem.

Crie uma Cloud TPU

A prática recomendada é aprovisionar Cloud TPUs v5es como recursos em fila usando o comando queued-resource create. Para mais informações, consulte o artigo Faça a gestão dos recursos em fila.

Também pode usar a API Create Node (gcloud compute tpus tpu-vm create) para aprovisionar TPUs em nuvem v5es. Para mais informações, consulte o artigo Faça a gestão dos recursos de TPU.

Para mais informações sobre as configurações v5e disponíveis para preparação, consulte o artigo Tipos de TPU na nuvem v5e para preparação.

Configuração da framework

Esta secção descreve o processo de configuração geral para a preparação de modelos personalizados com o JAX ou o PyTorch com a TPU v5e.

Para ver instruções de configuração da inferência, consulte o artigo Introdução à inferência da v5e.

Defina algumas variáveis de ambiente:

export PROJECT_ID=your_project_ID
export ACCELERATOR_TYPE=v5litepod-16
export ZONE=us-west4-a
export TPU_NAME=your_tpu_name
export QUEUED_RESOURCE_ID=your_queued_resource_id

Configuração do JAX

Se tiver formas de fatia com mais de 8 chips, terá várias VMs numa fatia. Neste caso, tem de usar a flag --worker=all para executar a instalação em todas as VMs de TPU num único passo sem usar SSH para iniciar sessão em cada uma separadamente:

gcloud compute tpus tpu-vm ssh ${TPU_NAME}  \
   --project=${PROJECT_ID} \
   --zone=${ZONE} \
   --worker=all \
   --command='pip install -U "jax[tpu]" -f https://storage.googleapis.com/jax-releases/libtpu_releases.html'

Descrições das flags de comando

  • TPU_NAME: O ID de texto atribuído pelo utilizador da TPU que é criado quando o pedido de recurso em fila é atribuído.
  • PROJECT_ID: Google Cloud Project Name. Use um projeto existente ou crie um novo em Configure o seu Google Cloud projeto
  • ZONE: consulte o documento Regiões e zonas da TPU para ver as zonas suportadas.
  • worker: a VM da TPU que tem acesso às TPUs subjacentes.

Pode executar o seguinte comando para verificar o número de dispositivos (os resultados apresentados aqui foram produzidos com uma fatia v5litepod-16). Este código testa se tudo está instalado corretamente, verificando se o JAX vê os TensorCores do Cloud TPU e consegue executar operações básicas:

gcloud compute tpus tpu-vm ssh ${TPU_NAME} \
   --project=${PROJECT_ID} \
   --zone=${ZONE} \
   --worker=all \
   --command='python3 -c "import jax; print(jax.device_count()); print(jax.local_device_count())"'

O resultado será semelhante ao seguinte:

SSH: Attempting to connect to worker 0...
SSH: Attempting to connect to worker 1...
SSH: Attempting to connect to worker 2...
SSH: Attempting to connect to worker 3...
16
4
16
4
16
4
16
4

jax.device_count() mostra o número total de chips na fatia especificada. jax.local_device_count() indica a contagem de chips acessíveis por uma única VM nesta fatia.

# Check the number of chips in the given slice by summing the count of chips
# from all VMs through the
# jax.local_device_count() API call.
gcloud compute tpus tpu-vm ssh ${TPU_NAME} \
   --project=${PROJECT_ID} \
   --zone=${ZONE} \
   --worker=all \
   --command='python3 -c "import jax; xs=jax.numpy.ones(jax.local_device_count()); print(jax.pmap(lambda x: jax.lax.psum(x, \"i\"), axis_name=\"i\")(xs))"'

O resultado será semelhante ao seguinte:

SSH: Attempting to connect to worker 0...
SSH: Attempting to connect to worker 1...
SSH: Attempting to connect to worker 2...
SSH: Attempting to connect to worker 3...
[16. 16. 16. 16.]
[16. 16. 16. 16.]
[16. 16. 16. 16.]
[16. 16. 16. 16.]

Experimente os tutoriais do JAX neste documento para começar a usar o treino v5e com o JAX.

Configuração do PyTorch

Tenha em atenção que a v5e só suporta o tempo de execução PJRT e o PyTorch 2.1+ vai usar o PJRT como o tempo de execução predefinido para todas as versões de TPU.

Esta secção descreve como começar a usar o PJRT na v5e com o PyTorch/XLA com comandos para todos os trabalhadores.

Instale dependências

gcloud compute tpus tpu-vm ssh ${TPU_NAME}  \
   --project=${PROJECT_ID} \
   --zone=${ZONE} \
   --worker=all \
   --command='
      sudo apt-get update -y
      sudo apt-get install libomp5 -y
      pip install mkl mkl-include
      pip install tf-nightly tb-nightly tbp-nightly
      pip install numpy
      sudo apt-get install libopenblas-dev -y
      pip install torch~=PYTORCH_VERSION torchvision torch_xla[tpu]~=PYTORCH_VERSION -f https://storage.googleapis.com/libtpu-releases/index.html -f https://storage.googleapis.com/libtpu-wheels/index.html'

Substitua PYTORCH_VERSION pela versão do PyTorch que quer usar. PYTORCH_VERSION é usado para especificar a mesma versão para o PyTorch/XLA. Recomendamos a versão 2.6.0.

Para mais informações sobre as versões do PyTorch e do PyTorch/XLA, consulte os artigos PyTorch – Começar e Lançamentos do PyTorch/XLA.

Para mais informações sobre a instalação do PyTorch/XLA, consulte o artigo Instalação do PyTorch/XLA.

Se receber um erro ao instalar as rodas para torch, torch_xla ou torchvision, como pkg_resources.extern.packaging.requirements.InvalidRequirement: Expected end or semicolon (after name and no valid version specifier) torch==nightly+20230222, use o seguinte comando para mudar para uma versão anterior: