Esegui un carico di lavoro batch con Pathways

Ai fini di questo documento, i carichi di lavoro batch sono definiti come carichi di lavoro JAX che vengono eseguiti fino al completamento e vengono implementati nello stesso cluster GKE del cluster Pathways, in particolare insieme ai componenti del controller Pathways (server proxy IFRT e gestore risorse Pathways). Il completamento del carico di lavoro JAX termina i componenti del cluster Pathways. Questa guida utilizza un carico di lavoro di addestramento JAX per dimostrarlo.

Prima di iniziare

Assicurati di avere:

Creare un'immagine di addestramento utilizzando Maxtext

MaxText è un progetto open source di modello linguistico di grandi dimensioni (LLM) sviluppato da Google. È scritto in JAX e progettato per essere altamente performante e scalabile, con un'esecuzione efficiente su TPU e GPU di Google Cloud.

Per creare un'immagine Docker MaxText utilizzando l'ultima versione stabile di JAX dal repository GitHub OSS, esegui questo comando:

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.

Questo comando esegue il push dell'immagine Kubernetes MaxText su gcr.io/$PROJECT_ID/${USER}_runner. Puoi utilizzare questa immagine Docker per eseguire l'addestramento sulle TPU utilizzando il backend Pathways.

Esegui un carico di lavoro batch con Cluster Toolkit

Invia l'immagine Docker MaxText predefinita utilizzando il comando 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"

Per saperne di più sulle opzioni di invio dei job, consulta la Guida all'invio dei job di Cluster Toolkit.

Sostituisci quanto segue:

  • WORKLOAD: un nome univoco per identificare il workload; a causa dei limiti delle etichette DNS, deve contenere al massimo 22 caratteri
  • CLUSTER: il nome del tuo cluster GKE
  • WORKLOAD_NODEPOOL_COUNT: il numero di node pool di sezioni TPU
  • COMPUTE_TYPE: il tipo di macchina TPU (ad esempio ct6e-standard-4t). Per maggiori informazioni sui tipi di TPU supportati per ogni versione della TPU, consulta Versioni della TPU.
  • TOPOLOGY: la topologia di posizionamento TPU (ad esempio, 2x4)
  • PROJECT_ID: il tuo Google Cloud ID progetto
  • ZONE: la zona in cui prevedi di eseguire il tuo workload
  • USER: il tuo Google Cloud ID utente
  • BUCKET_NAME: il bucket Cloud Storage per gli output
  • RUN_NAME: un nome assegnato dall'utente per identificare l'esecuzione del workflow

Segui l'avanzamento del tuo workload utilizzando il comando 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

Per annullare il carico di lavoro prima del completamento, utilizza il comando gcluster job cancel:

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

Passaggi successivi