Profiler JAX sur des GPU avec XProf et les diagnostics ML

Pour optimiser les modèles JAX à grande échelle sur les GPU, il est nécessaire d'avoir une visibilité approfondie sur les goulots d'étranglement des performances. Ce guide fournit un workflow complet de bout en bout pour exécuter et profiler des charges de travail JAX sur des GPU (tels que NVIDIA L4) à l'aide de Google Cloud ML Diagnostics et XProf. En tirant parti de ces outils, vous pouvez identifier les opérations inefficaces, optimiser l'utilisation des ressources de calcul et accélérer vos entraînements.

En suivant ce guide, vous apprendrez à :

  1. Instrumentez une simple boucle d'entraînement JAX pour le profilage.
  2. Conteneurisez la charge de travail avec la prise en charge CUDA appropriée.
  3. Déployez la charge de travail sur Google Kubernetes Engine (GKE) à l'aide de JobSet.
  4. Capturez et visualisez les profils de performances de manière dynamique.

Prérequis

À vérifier avant de commencer :

  • Un projet Google Cloud avec facturation activée.
  • Un cluster GKE compatible avec les GPU (par exemple, NVIDIA L4).
  • Un bucket Google Cloud Storage (GCS) pour stocker les profils.
  • Les CLI gcloud et kubectl sont installées et configurées.
  • Workload Identity configuré pour votre cluster GKE afin d'accéder à GCS.

Étape 1 : Instrumenter la charge de travail JAX

Commencez par créer un script d'entraînement JAX (par exemple, train.py). Nous utilisons le SDK google-cloud-mldiagnostics pour interagir avec l'infrastructure de profilage gérée.

[!WARNING] Le script ci-dessous inclut une boucle infinie pour maintenir le GPU occupé lors des démonstrations de profilage à la demande. N'oubliez pas d'arrêter manuellement le job ou de supprimer les ressources GKE une fois que vous avez terminé pour éviter des coûts de facturation inutiles.

import logging
import os
import time
from google_cloud_mldiagnostics import machinelearning_run
from google_cloud_mldiagnostics import xprof
import jax
import jax.numpy as jnp
import numpy as np

logging.basicConfig(level=logging.INFO)

def main():
    logging.info("Starting JAX training job...")

    # Coordinates multihost collective operations and healthchecks
    jax.distributed.initialize()

    logging.info(
        f"JAX initialized: process_index={jax.process_index()}, "
        f"process_count={jax.process_count()}"
    )

    # Syncs metadata with the mldiag hook & launches reverse proxy daemons
    machinelearning_run(
        name=f"jax-gpu-run-{int(time.time())}",
        configs={"learning_rate": 1e-5, "batch_size": 8192},
        project=os.environ.get("PROJECT_ID", "<your-project-id>"),
        region=os.environ.get("REGION", "us-central1"),
        gcs_path=os.environ.get("GCS_BUCKET", "gs://<your-gcs-bucket>"),
        on_demand_xprof=True,
    )

    key = jax.random.PRNGKey(0)
    size = 4096
    matrix = jax.random.normal(key, (size, size), dtype=jnp.float32)

    def train_step(x):
        return jnp.dot(x, x)

    train_step = jax.jit(train_step)

    # Triggers XLA compilation ahead of tracing steps so compilation overhead isn't profiled
    matrix = train_step(matrix)
    matrix.block_until_ready() # Wait for compilation to complete.

    prof = xprof()
    prof.start(session_id="warmup_phase")

    for _ in range(5):
        matrix = train_step(matrix)
        matrix.block_until_ready()

    prof.stop()
    logging.info("Programmatic profile capture complete.")

    logging.info("Entering training loop. Ready for on-demand profiling...")
    try:
        while True:
            # Continuously pump steps keeping GPUs occupied for on-demand capture triggers
            matrix = train_step(matrix)
            matrix.block_until_ready() # Ensure GPU work completes before next step.
            time.sleep(0.5)
    except KeyboardInterrupt:
        logging.info("Training loop interrupted.")

if __name__ == "__main__":
    main()

Étape 2 : Conteneurisation (Dockerfile)

Créez un Dockerfile pour regrouper votre script JAX avec les dépendances CUDA requises et le SDK ML Diagnostics.

# Use an official NVIDIA CUDA base image compatible with JAX
FROM nvidia/cuda:13.2.1-cudnn-devel-ubuntu24.04

# Install Python, venv, and other OS dependencies
RUN apt-get update && apt-get install -y \
    python3-pip \
    python3-venv \
    python3-dev \
    git \
    curl \
    && rm -rf /var/lib/apt/lists/*

# Set up a virtual environment and update PATH to use it implicitly
ENV VIRTUAL_ENV=/opt/venv
RUN python3 -m venv $VIRTUAL_ENV
ENV PATH="$VIRTUAL_ENV/bin:$PATH"

# At this point, pip and python implicitly map to the virtual env!
# No need for --break-system-packages.

# Upgrade pip inside the venv
RUN pip install --upgrade pip

# Install JAX with CUDA support
RUN pip install --upgrade "jax[cuda13]" -f https://storage.googleapis.com/jax-releases/jax_cuda_releases.html

# Install ML Diagnostics SDK and XProf tools
RUN pip install --no-cache-dir \
    google-cloud-mldiagnostics \
    xprof-nightly

WORKDIR /app
COPY train.py .

CMD ["python3", "train.py"]

Créez et transférez l'image vers Artifact Registry :

docker build -t us-central1-docker.pkg.dev/<project-id>/<repo>/jax-gpu-workload:latest .
docker push us-central1-docker.pkg.dev/<project-id>/<repo>/jax-gpu-workload:latest

Étape 3 : Déploiement (manifeste Kubernetes)

Déployez la charge de travail à l'aide d'un JobSet GKE ou d'un job standard. Pour permettre à la plate-forme de diagnostics ML d'injecter des métadonnées et de router les demandes de profil, appliquez le libellé managed-mldiagnostics-gke: "true". Pour en savoir plus sur la configuration de GKE pour les diagnostics ML, consultez le guide de configuration de GKE officiel.

apiVersion: jobset.x-k8s.io/v1alpha2
kind: JobSet
metadata:
  name: jax-gpu-job
  namespace: ai-workloads
  labels:
    managed-mldiagnostics-gke: "true"
spec:
  replicatedJobs:
  - name: gpu-nodes
    replicas: 1
    template:
      spec:
        parallelism: 1
        completions: 1
        backoffLimit: 0
        template:
          metadata:
            labels:
              managed-mldiagnostics-gke: "true"
          spec:
            # Must match the GKE Service Account with Workload Identity permissions
            serviceAccountName: <your-service-account>
            hostNetwork: true
            dnsPolicy: ClusterFirstWithHostNet
            nodeSelector:
              cloud.google.com/gke-accelerator: nvidia-l4 # Or other GPU
            containers:
            - name: workload
              image: us-central1-docker.pkg.dev/<project-id>/<repo>/jax-gpu-workload:latest
              imagePullPolicy: Always
              # Expose ports required for profile daemons
              ports:
              - containerPort: 8471 # JAX distributed coordinator port
              - containerPort: 8080 # ML Diagnostics agent/proxy port
              - containerPort: 9999 # XProf server port for on-demand profiling
              resources:
                limits:
                  nvidia.com/gpu: 1

Appliquez le fichier manifeste :

kubectl apply -f deploy.yaml

Étape 4 : Capture et visualisation

Capture programmatique

Si vous avez inclus prof.start() / prof.stop() dans votre script, ces profils sont automatiquement importés dans votre bucket GCS sous le chemin d'accès suivant : gs://<your-gcs-bucket>/<run-name>/plugins/profile/<session-id>/.

Capture à la demande

Comme on_demand_xprof=True est défini dans machinelearning_run, vous pouvez capturer des profils de manière dynamique pendant l'exécution du job.

Pour obtenir des instructions détaillées sur l'utilisation de l'interface utilisateur TensorBoard afin de déclencher des profils à la demande, de sélectionner des pods spécifiques et d'afficher les traces capturées, veuillez consulter la documentation publique officielle : Diagnostics ML Google Cloud : capture de profil à la demande.

Vous pouvez également capturer des profils à l'aide de gcloud CLI, comme décrit dans le guide de la CLI ML Diagnostics.

Cette documentation publique s'applique aux charges de travail TPU et GPU gérées par ML Diagnostics.