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:
- Instrumentieren Sie eine einfache JAX-Trainingsschleife für die Profilerstellung.
- Containerisieren Sie die Arbeitslast mit der entsprechenden CUDA-Unterstützung.
- Stellen Sie die Arbeitslast mit JobSet in Google Kubernetes Engine (GKE) bereit.
- 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- undkubectl-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.