Entraîner un modèle à l'aide de TPU v5e
Avec une empreinte de 256 puces par pod, les TPU v5e sont optimisés pour être un produit à forte valeur ajoutée pour l'entraînement, l'ajustement et la mise en service des transformateurs, du texte à l'image et des réseaux de neurones convolutifs (CNN). Pour en savoir plus sur l'utilisation de Cloud TPU v5e pour la mise en service, consultez Inférence avec v5e.
Pour en savoir plus sur le matériel et les configurations Cloud TPU v5e, consultez TPU v5e.
Commencer
Les sections suivantes décrivent comment commencer à utiliser les TPU v5e.
Quota de requêtes
Vous avez besoin d'un quota pour utiliser des TPU v5e pour l'entraînement. Il existe différents types de quotas pour les TPU à la demande, les TPU réservés et les VM Spot TPU. Des quotas distincts sont requis si vous utilisez votre TPU v5e pour l'inférence. Pour en savoir plus sur les quotas, consultez Quotas. Pour demander un quota TPU v5e, contactez le service commercial Cloud.
Créer un compte et un projet Google Cloud
Vous avez besoin d'un compte Google Cloud et d'un projet pour utiliser Cloud TPU. Pour en savoir plus, consultez Configurer un environnement Cloud TPU.
Créer une instance Cloud TPU
La bonne pratique consiste à provisionner les Cloud TPU v5e en tant que ressources mises en file d'attente à l'aide de la commande queued-resource create. Pour en savoir plus, consultez Gérer les ressources mises en file d'attente.
Vous pouvez également utiliser l'API Create Node (gcloud compute tpus tpu-vm create) pour provisionner des Cloud TPU v5e. Pour en savoir plus, consultez Gérer les ressources TPU.
Pour en savoir plus sur les configurations v5e disponibles pour l'entraînement, consultez Types de Cloud TPU v5e pour l'entraînement.
Configurer le framework
Cette section décrit la procédure de configuration générale pour l'entraînement de modèles personnalisés à l'aide de JAX ou PyTorch avec des TPU v5e.
Pour obtenir des instructions de configuration de l'inférence, consultez Présentation de l'inférence v5e.
Définir certaines variables d'environnement :
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
Configurer pour JAX
Si vous avez des formes de tranche supérieures à huit puces, vous aurez plusieurs VM dans une même tranche. Dans ce cas, vous devez utiliser le flag --worker=all pour exécuter l'installation sur toutes les VM TPU en une seule étape, sans utiliser SSH pour vous connecter à chacune d'elles séparément :
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'
Description des flags de commande
| Variable | Description |
| TPU_NAME | ID de texte attribué par l'utilisateur au TPU créé lorsque la demande de ressource en file d'attente est allouée. |
| PROJECT_ID | Nom du projetGoogle Cloud Utilisez un projet existant ou créez-en un dans Configurer votre projet Google Cloud . |
| ZONE | Pour connaître les zones compatibles, consultez le document Régions et zones TPU. |
| nœud de calcul | VM TPU ayant accès aux TPU sous-jacents. |
Vous pouvez exécuter la commande suivante pour vérifier le nombre d'appareils (les sorties affichées ici ont été produites avec une tranche v5litepod-16). Ce code teste que tout est correctement installé en vérifiant que JAX voit les TensorCores Cloud TPU et qu'il peut exécuter des opérations de base :
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())"'
Le résultat doit ressembler à ce qui suit :
SSH: Attempting to