Generación de perfiles de JAX en GPUs con XProf y ML Diagnostics

La optimización de modelos JAX a gran escala en GPUs requiere una visibilidad profunda de los cuellos de botella del rendimiento. En esta guía, se proporciona un flujo de trabajo integral de extremo a extremo (E2E) para ejecutar y generar perfiles de cargas de trabajo de JAX en GPUs (como NVIDIA L4) con Google Cloud ML Diagnostics y XProf. Con estas herramientas, puedes identificar operaciones ineficientes, optimizar el uso de los recursos de procesamiento y acelerar tus ejecuciones de entrenamiento.

Si sigues esta guía, aprenderás a hacer lo siguiente:

  1. Instrumenta un bucle de entrenamiento de JAX simple para la creación de perfiles.
  2. Alojamiento de la carga de trabajo en contenedores con la compatibilidad de CUDA adecuada
  3. Implementa la carga de trabajo en Google Kubernetes Engine (GKE) con JobSet.
  4. Captura y visualiza perfiles de rendimiento de forma dinámica.

Requisitos previos

Antes de comenzar, asegúrate de contar con los siguientes aspectos:

  • Un proyecto de Google Cloud con facturación habilitada.
  • Un clúster de GKE con compatibilidad con GPU (p. ej., NVIDIA L4)
  • Un bucket de Google Cloud Storage (GCS) para almacenar perfiles
  • Las CLIs de gcloud y kubectl deben estar instaladas y configuradas.
  • Workload Identity configurada para que tu clúster de GKE acceda a GCS

Paso 1: Instrumenta la carga de trabajo de JAX

Primero, crea una secuencia de comandos de entrenamiento de JAX (p.ej., train.py). Usamos el SDK de google-cloud-mldiagnostics para interactuar con la infraestructura de generación de perfiles administrada.

[!WARNING] La siguiente secuencia de comandos incluye un bucle infinito para mantener la GPU ocupada durante las demostraciones de generación de perfiles a pedido. Recuerda detener el trabajo o borrar los recursos de GKE de forma manual cuando termines para evitar costos de facturación innecesarios.

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

Paso 2: Creación de contenedores (Dockerfile)

Crea un Dockerfile para empaquetar tu secuencia de comandos de JAX con las dependencias de CUDA requeridas y el SDK de 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"]

Compila y envía la imagen a 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

Paso 3: Implementación (manifiesto de Kubernetes)

Implementa la carga de trabajo con un JobSet de GKE o un Job estándar. Para permitir que la plataforma de ML Diagnostics inserte metadatos y enrute solicitudes de perfil, aplica la etiqueta managed-mldiagnostics-gke: "true". Para obtener más detalles sobre la configuración de GKE para ML Diagnostics, consulta la Guía de configuración de GKE oficial.

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

Aplica el manifiesto

kubectl apply -f deploy.yaml

Paso 4: Captura y visualización

Captura programática

Si incluiste prof.start() / prof.stop() en tu secuencia de comandos, esos perfiles se subirán automáticamente a tu bucket de GCS en la siguiente ruta de acceso: gs://<your-gcs-bucket>/<run-name>/plugins/profile/<session-id>/

Captura bajo demanda

Como on_demand_xprof=True se configura en machinelearning_run, puedes capturar perfiles de forma dinámica mientras se ejecuta el trabajo.

Para obtener instrucciones detalladas sobre cómo usar la IU de TensorBoard para activar perfiles a pedido, seleccionar Pods específicos y ver los registros capturados, consulta la documentación pública oficial: Google Cloud ML Diagnostics: Captura de perfiles a pedido.

También puedes capturar perfiles con la gcloud CLI, como se describe en la Guía de la CLI de ML Diagnostics.

Esta documentación pública se aplica a las cargas de trabajo de TPU y GPU que administra ML Diagnostics.