Profilazione di JAX su GPU con XProf e ML Diagnostics

L'ottimizzazione di modelli JAX su larga scala sulle GPU richiede una visibilità approfondita dei colli di bottiglia delle prestazioni. Questa guida fornisce un flusso di lavoro end-to-end (E2E) completo per l'esecuzione e la profilazione dei workload JAX su GPU (come NVIDIA L4) utilizzando Google Cloud ML Diagnostics e XProf. Sfruttando questi strumenti, puoi identificare operazioni inefficienti, ottimizzare l'utilizzo delle risorse di calcolo e accelerare le esecuzioni di addestramento.

Seguendo questa guida, imparerai a:

  1. Strumenta un semplice ciclo di addestramento JAX per la profilazione.
  2. Containerizza il workload con il supporto CUDA appropriato.
  3. Esegui il deployment del workload su Google Kubernetes Engine (GKE) utilizzando JobSet.
  4. Acquisire e visualizzare dinamicamente i profili delle prestazioni.

Prerequisiti

Prima di iniziare, assicurati di avere:

  • Un progetto Google Cloud con la fatturazione abilitata.
  • Un cluster GKE con supporto GPU (ad es. NVIDIA L4).
  • Un bucket Google Cloud Storage (GCS) per archiviare i profili.
  • CLI gcloud e kubectl installate e configurate.
  • Workload Identity configurata per il cluster GKE per accedere a GCS.

Passaggio 1: strumentazione del workload JAX

Innanzitutto, crea uno script di addestramento JAX (ad es. train.py). Utilizziamo l'SDK google-cloud-mldiagnostics per interagire con l'infrastruttura di profilazione gestita.

[!WARNING] Lo script riportato di seguito include un ciclo infinito per mantenere la GPU occupata per le dimostrazioni di profilazione on demand. Ricorda di arrestare manualmente il job o eliminare le risorse GKE al termine per evitare costi di fatturazione inutili.

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()

Passaggio 2: containerizzazione (Dockerfile)

Crea un Dockerfile per pacchettizzare lo script JAX con le dipendenze CUDA richieste e l'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"]

Crea l'immagine ed eseguine il push in 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

Passaggio 3: deployment (manifest Kubernetes)

Esegui il deployment del workload utilizzando un JobSet GKE o un job standard. Per consentire alla piattaforma ML Diagnostics di inserire metadati e indirizzare le richieste di profili, applica l'etichetta managed-mldiagnostics-gke: "true". Per saperne di più sulla configurazione di GKE per ML Diagnostics, consulta la Guida alla configurazione di GKE ufficiale.

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

Applica il manifest:

kubectl apply -f deploy.yaml

Passaggio 4: acquisizione e visualizzazione

Acquisizione programmatica

Se hai incluso prof.start() / prof.stop() nello script, questi profili vengono caricati automaticamente nel bucket GCS nel percorso: gs://<your-gcs-bucket>/<run-name>/plugins/profile/<session-id>/

Acquisizione on demand

Poiché on_demand_xprof=True è impostato in machinelearning_run, puoi acquisire i profili in modo dinamico durante l'esecuzione del job.

Per istruzioni dettagliate su come utilizzare la UI di TensorBoard per attivare profili on demand, selezionare pod specifici e visualizzare le tracce acquisite, consulta la documentazione pubblica ufficiale: Google Cloud ML Diagnostics - On-demand profile capture.

Puoi anche acquisire profili utilizzando gcloud CLI come descritto nella guida ML Diagnostics CLI.

Questa documentazione pubblica si applica sia ai workload TPU che GPU gestiti da ML Diagnostics.