JAX auf GPUs mit XProf und ML Diagnostics profilieren

Für die Optimierung von JAX-Modellen im großen Maßstab auf GPUs ist eine detaillierte Analyse von Leistungsengpässen erforderlich. In dieser Anleitung wird ein umfassender End-to-End-Workflow (E2E) zum Ausführen und Profilieren von JAX-Arbeitslasten auf GPUs (z. B. NVIDIA L4) mit Google Cloud ML Diagnostics und XProf beschrieben. Mithilfe dieser Tools können Sie ineffiziente Vorgänge identifizieren, die Nutzung von Rechenressourcen optimieren und Ihre Trainingsläufe beschleunigen.

In diesem Leitfaden erfahren Sie, wie Sie:

  1. Instrumentieren Sie eine einfache JAX-Trainingsschleife für die Profilerstellung.
  2. Containerisieren Sie die Arbeitslast mit der entsprechenden CUDA-Unterstützung.
  3. Stellen Sie die Arbeitslast mit JobSet in Google Kubernetes Engine (GKE) bereit.
  4. Leistungsprofile dynamisch erfassen und visualisieren

Vorbereitung

Prüfen Sie zuerst, ob Sie Folgendes haben:

  • Google Cloud-Projekt mit aktivierter Abrechnungsfunktion.
  • Ein GKE-Cluster mit GPU-Unterstützung (z.B. NVIDIA L4).
  • Ein Google Cloud Storage-Bucket (GCS) zum Speichern von Profilen.
  • Die gcloud- und kubectl-Befehlszeilen sind installiert und konfiguriert.
  • Workload Identity für Ihren GKE-Cluster für den Zugriff auf GCS konfiguriert.

Schritt 1: JAX-Arbeitslast instrumentieren

Erstellen Sie zuerst ein JAX-Trainingsskript (z.B. train.py). Wir verwenden das google-cloud-mldiagnostics SDK, um mit der verwalteten Profilerstellungsinfrastruktur zu interagieren.

[!WARNING] Das folgende Skript enthält eine Endlosschleife, um die GPU für On-Demand-Profiling-Demonstrationen zu beschäftigen. Denken Sie daran, den Job manuell zu beenden oder die GKE-Ressourcen zu löschen, wenn Sie fertig sind, um unnötige Abrechnungskosten zu vermeiden.

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

Schritt 2: Containerisierung (Dockerfile)

Erstellen Sie ein Dockerfile, um Ihr JAX-Skript mit den erforderlichen CUDA-Abhängigkeiten und dem ML Diagnostics SDK zu verpacken.

# 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"]

Erstellen Sie das Image und übertragen Sie es per Push in die 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

Schritt 3: Bereitstellung (Kubernetes-Manifest)

Stellen Sie die Arbeitslast mit einem GKE-JobSet oder einem Standard-Job bereit. Damit die ML-Diagnoseplattform Metadaten einfügen und Profilanfragen weiterleiten kann, wenden Sie das Label managed-mldiagnostics-gke: "true" an. Weitere Informationen zum Konfigurieren von GKE für ML Diagnostics finden Sie im offiziellen GKE-Einrichtungsleitfaden.

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

Wenden Sie das Manifest an:

kubectl apply -f deploy.yaml

Schritt 4: Erfassung und Visualisierung

Programmatische Erfassung

Wenn Sie prof.start() / prof.stop() in Ihr Script aufgenommen haben, werden diese Profile automatisch in Ihren GCS-Bucket unter dem Pfad gs://<your-gcs-bucket>/<run-name>/plugins/profile/<session-id>/ hochgeladen.

On-Demand-Erfassung

Da on_demand_xprof=True in machinelearning_run festgelegt ist, können Sie Profile dynamisch erfassen, während der Job ausgeführt wird.

Eine detaillierte Anleitung zur Verwendung der TensorBoard-Benutzeroberfläche zum Auslösen von On-Demand-Profilen, zum Auswählen bestimmter Pods und zum Ansehen der erfassten Traces finden Sie in der offiziellen öffentlichen Dokumentation: Google Cloud – ML Diagnostics – On-Demand-Profilerstellung.

Sie können Profile auch mit der gcloud CLI erfassen, wie im CLI-Leitfaden für ML-Diagnosen beschrieben.

Diese öffentliche Dokumentation gilt sowohl für TPU- als auch für GPU-Arbeitslasten, die von ML Diagnostics verwaltet werden.