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:
- Strumenta un semplice ciclo di addestramento JAX per la profilazione.
- Containerizza il workload con il supporto CUDA appropriato.
- Esegui il deployment del workload su Google Kubernetes Engine (GKE) utilizzando JobSet.
- 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
gcloudekubectlinstallate 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.