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 projetoZONE: 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: