Exécuter une charge de travail par lot avec Pathways

Dans ce document, les charges de travail par lot sont définies comme des charges de travail JAX qui s'exécutent jusqu'à la fin et sont déployées dans le même cluster GKE que le cluster Pathways, plus précisément aux côtés des composants du contrôleur Pathways (serveur proxy IFRT et gestionnaire de ressources Pathways). Une fois la charge de travail JAX terminée, les composants du cluster Pathways sont arrêtés. Ce guide utilise une charge de travail d'entraînement JAX pour illustrer ce point.

Avant de commencer

Vérifiez que vous disposez bien des éléments suivants :

Créer une image d'entraînement à l'aide de MaxText

MaxText est un projet de grand modèle de langage (LLM) Open Source développé par Google. Il est écrit en JAX et conçu pour être très performant et évolutif, et s'exécute efficacement sur les TPU et GPU Google Cloud.

Pour créer une image Docker MaxText à l'aide de la dernière version stable de JAX à partir du dépôt GitHub OSS, exécutez la commande suivante :

git clone https://github.com/AI-Hypercomputer/maxtext
cd maxtext/dependencies/scripts
gcloud config set project PROJECT_ID
bash ./docker_build_dependency_image.sh MODE=stable
gcloud auth configure-docker
bash ./docker_upload_runner.sh CLOUD_IMAGE_NAME=USER_runner # This script needs bash version >= 4.2 to execute.

Cette commande transfère l'image Kubernetes MaxText vers gcr.io/$PROJECT_ID/${USER}_runner. Vous pouvez utiliser cette image Docker pour exécuter l'entraînement sur des TPU à l'aide du backend Pathways.

Exécuter une charge de travail par lot avec Cluster Toolkit

Envoyez l'image Docker MaxText préconfigurée à l'aide de la commande gcluster job submit :

gcluster job submit \
    --pathways \
    --pathways-gcs-location="gs://BUCKET_NAME/pathways-artifacts" \
    --name=WORKLOAD \
    --cluster=CLUSTER \
    --project=PROJECT_ID \
    --location=ZONE \
    --num-slices=WORKLOAD_NODEPOOL_COUNT \
    --compute-type=COMPUTE_TYPE \
    --topology=TOPOLOGY \
    --image="gcr.io/PROJECT_ID/USER_runner" \
    --command="python3 -m MaxText.train /deps/src/MaxText/configs/base.yml base_output_directory=gs://BUCKET_NAME per_device_batch_size=1 enable_checkpointing=false remat_policy=full global_parameter_scale=1 steps=20 max_target_length=2048 use_iota_embed=true reuse_example_batch=1 dataset_type=synthetic attention=flash gcs_metrics=True enable_single_controller=True run_name=RUN_NAME-pathways-job"

Pour en savoir plus sur les options d'envoi de tâches, consultez le guide d'envoi de tâches Cluster Toolkit.

Remplacez les éléments suivants :

  • WORKLOAD : nom unique permettant d'identifier votre charge de travail. En raison des limites des libellés DNS, il doit comporter 22 caractères ou moins.
  • CLUSTER : nom de votre cluster GKE
  • WORKLOAD_NODEPOOL_COUNT : nombre de pools de nœuds de tranche TPU
  • COMPUTE_TYPE : type de machine TPU (par exemple, ct6e-standard-4t). Pour en savoir plus sur les types de TPU compatibles avec chaque version de TPU, consultez Versions de TPU.
  • TOPOLOGY : topologie de placement des TPU (par exemple, 2x4)
  • PROJECT_ID : ID de votre projet Google Cloud
  • ZONE : zone dans laquelle vous prévoyez d'exécuter votre charge de travail
  • USER : votre ID utilisateur Google Cloud
  • BUCKET_NAME : bucket Cloud Storage pour les sorties
  • RUN_NAME : nom attribué par l'utilisateur pour identifier l'exécution du workflow

Suivez la progression de votre charge de travail à l'aide de la commande gcluster job logs :

gcluster job logs WORKLOAD \
    --cluster=CLUSTER \
    --project=PROJECT_ID \
    --location=ZONE \
    --main-only=false
completed step: 1, seconds: 0.484, TFLOP/s/device: 87.349, Tokens/s/device: 2117.382, total_weights: 2945, loss: 10.888
completed step: 2, seconds: 0.407, TFLOP/s/device: 103.699, Tokens/s/device: 2513.735, total_weights: 3253, loss: 9.697
completed step: 3, seconds: 0.248, TFLOP/s/device: 170.300, Tokens/s/device: 4128.167, total_weights: 3154, loss: 9.641
completed step: 4, seconds: 0.216, TFLOP/s/device: 195.122, Tokens/s/device: 4729.880, total_weights: 3119, loss: 9.547
completed step: 5, seconds: 0.272, TFLOP/s/device: 155.298, Tokens/s/device: 3764.512, total_weights: 2837, loss: 10.179
completed step: 6, seconds: 0.472, TFLOP/s/device: 89.489, Tokens/s/device: 2169.266, total_weights: 3069, loss: 9.776

Pour annuler la charge de travail avant qu'elle ne soit terminée, utilisez la commande gcluster job cancel :

gcluster job cancel WORKLOAD --cluster=CLUSTER --project=PROJECT_ID --location=ZONE

Étapes suivantes