יצירת פרופיל של JAX במעבדי GPU באמצעות XProf ו-ML Diagnostics

כדי לבצע אופטימיזציה של מודלים גדולים של JAX ב-GPU, צריך לקבל תובנות מעמיקות לגבי צווארי בקבוק בביצועים. במדריך הזה מוסבר איך להריץ ולנתח פרופילים של עומסי עבודה של JAX ב-GPU (כמו NVIDIA L4) באמצעות Google Cloud ML Diagnostics ו-XProf. הכלים האלה מאפשרים לכם לזהות פעולות לא יעילות, לייעל את השימוש במשאבי מחשוב ולהאיץ את תהליכי האימון.

במדריך הזה תלמדו איך:

  1. לבצע אינסטרומנטציה ללולאת אימון פשוטה של JAX לצורך יצירת פרופיל.
  2. מכניסים את עומס העבודה לקונטיינר עם תמיכה מתאימה ב-CUDA.
  3. פריסת עומס העבודה ב-Google Kubernetes Engine‏ (GKE) באמצעות JobSet.
  4. תיעוד והמחשה של פרופילי ביצועים באופן דינמי.

דרישות מוקדמות

לפני שמתחילים, חשוב לוודא שיש לכם:

  • פרויקט ב-Google Cloud שהחיוב בו מופעל.
  • אשכול GKE עם תמיכה ב-GPU (לדוגמה, NVIDIA L4).
  • קטגוריה של Google Cloud Storage‏ (GCS) לאחסון פרופילים.
  • ה-CLI של gcloud ושל kubectl מותקנים ומוגדרים.
  • ‫Workload Identity מוגדר באשכול GKE כדי לגשת ל-GCS.

שלב 1: הטמעת JAX Workload

קודם כל, יוצרים סקריפט לאימון JAX (לדוגמה, train.py). משתמשים ב-google-cloud-mldiagnostics SDK כדי ליצור אינטראקציה עם התשתית המנוהלת של יצירת פרופילים.

‫[!WARNING] הסקריפט שלמטה כולל לולאה אינסופית כדי להעסיק את ה-GPU לצורך הדגמות של יצירת פרופילים לפי דרישה. כדי להימנע מחיובים מיותרים, חשוב לזכור להפסיק את העבודה באופן ידני או למחוק את משאבי GKE אחרי שמסיימים.

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

שלב 2: יצירת קונטיינר (Dockerfile)

יוצרים Dockerfile כדי לארוז את סקריפט ה-JAX עם יחסי התלות הנדרשים של CUDA ועם ערכת ה-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"]

יוצרים את התמונה ומעבירים אותה בדחיפה ל-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

שלב 3: פריסה (מניפסט של Kubernetes)

פורסים את עומס העבודה באמצעות JobSet או Job רגיל של GKE. כדי לאפשר לפלטפורמת האבחון של ML להוסיף מטא-נתונים ולנתב בקשות לפרופילים, צריך להחיל את התווית managed-mldiagnostics-gke: "true". פרטים נוספים על הגדרת GKE ל-ML Diagnostics מופיעים במדריך ההגדרה של 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

החלת המניפסט:

kubectl apply -f deploy.yaml

שלב 4: תיעוד והצגה חזותית

צילום פרוגרמטי

אם הוספתם prof.start() / prof.stop() לסקריפט, הפרופילים האלה יועלו אוטומטית לקטגוריית GCS בנתיב: gs://<your-gcs-bucket>/<run-name>/plugins/profile/<session-id>/

צילום על פי דרישה

מכיוון שהמשתנה on_demand_xprof=True מוגדר ב-machinelearning_run, אפשר ללכוד פרופילים באופן דינמי בזמן שהעבודה פועלת.

הוראות מפורטות לשימוש בממשק המשתמש של TensorBoard כדי להפעיל פרופילים על פי דרישה, לבחור פודים ספציפיים ולהציג את העקבות שתועדו זמינות במסמכים הרשמיים לציבור: Google Cloud ML Diagnostics - On-demand profile capture.

אפשר גם לצלם פרופילים באמצעות ה-CLI של gcloud, כמו שמתואר במדריך ל-CLI של ML Diagnostics.

התיעוד הזה שזמין לכולם רלוונטי לעומסי עבודה של TPU ו-GPU שמנוהלים על ידי ML Diagnostics.