Profilowanie JAX na GPU za pomocą XProf i ML Diagnostics

Optymalizacja modeli JAX na dużą skalę na procesorach GPU wymaga szczegółowego wglądu w wąskie gardła wydajności. Ten przewodnik zawiera kompleksowy przepływ pracy od początku do końca, który umożliwia uruchamianie i profilowanie zbiorów zadań JAX na procesorach graficznych (takich jak NVIDIA L4) przy użyciu narzędzi Google Cloud ML Diagnostics i XProf. Korzystając z tych narzędzi, możesz identyfikować nieefektywne operacje, optymalizować wykorzystanie zasobów obliczeniowych i przyspieszać trenowanie modeli.

Z tego przewodnika dowiesz się, jak:

  1. Instrumentuj prostą pętlę trenowania JAX na potrzeby profilowania.
  2. Skonteneryzuj zbiór zadań z odpowiednią obsługą CUDA.
  3. Wdróż zadanie w Google Kubernetes Engine (GKE) za pomocą JobSet.
  4. Dynamiczne rejestrowanie i wizualizowanie profili wydajności.

Wymagania wstępne

Zanim zaczniesz, upewnij się, że masz:

  • projekt Google Cloud z włączonymi płatnościami;
  • klaster GKE z obsługą GPU (np. NVIDIA L4);
  • Zasobnik Google Cloud Storage (GCS) do przechowywania profili.
  • Zainstalowane i skonfigurowane interfejsy wiersza poleceń gcloud i kubectl.
  • Workload Identity skonfigurowana dla klastra GKE w celu uzyskania dostępu do GCS.

Krok 1. Instrumentacja zbioru zadań JAX

Najpierw utwórz skrypt trenowania JAX (np. train.py). Do interakcji z zarządzaną infrastrukturą profilowania używamy pakietu SDK google-cloud-mldiagnostics.

[!WARNING] Poniższy skrypt zawiera pętlę nieskończoną, która utrzymuje GPU w stanie zajętości na potrzeby demonstracji profilowania na żądanie. Po zakończeniu pracy pamiętaj, aby ręcznie zatrzymać zadanie lub usunąć zasoby GKE, co pozwoli Ci uniknąć niepotrzebnych opłat.

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

Krok 2. Konteneryzacja (Dockerfile)

Utwórz Dockerfile, aby spakować skrypt JAX z wymaganymi zależnościami CUDA i pakietem 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"]

Utwórz obraz i przenieś go do 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

Krok 3. Wdrożenie (manifest Kubernetes)

Wdróż zbiór zadań za pomocą JobSet GKE lub standardowego zadania. Aby umożliwić platformie diagnostyki ML wstrzykiwanie metadanych i kierowanie żądań profili, zastosuj etykietę managed-mldiagnostics-gke: "true". Więcej informacji o konfigurowaniu GKE pod kątem diagnostyki ML znajdziesz w oficjalnym przewodniku po konfigurowaniu GKE.

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

Zastosuj plik manifestu:

kubectl apply -f deploy.yaml

Krok 4. Rejestrowanie i wizualizacja

Automatyzacja

Jeśli w skrypcie uwzględnisz prof.start() / prof.stop(), te profile zostaną automatycznie przesłane do zasobnika GCS w ścieżce:gs://<your-gcs-bucket>/<run-name>/plugins/profile/<session-id>/

Rejestrowanie na żądanie

Ponieważ w polu machinelearning_run ustawiono wartość on_demand_xprof=True, możesz dynamicznie rejestrować profile podczas działania zadania.

Szczegółowe instrukcje dotyczące korzystania z interfejsu TensorBoard do wywoływania profili na żądanie, wybierania konkretnych zasobników i wyświetlania przechwyconych śladów znajdziesz w oficjalnej dokumentacji publicznej: Diagnostyka ML w Google Cloud – przechwytywanie profili na żądanie.

Profile możesz też rejestrować za pomocą gcloud CLI, jak opisano w przewodniku po interfejsie ML Diagnostics CLI.

Ta publiczna dokumentacja dotyczy zarówno zbiorów zadań TPU, jak i GPU zarządzanych przez diagnostykę ML.