Membuat profil JAX di GPU dengan XProf dan Diagnostik ML

Mengoptimalkan model JAX skala besar di GPU memerlukan visibilitas mendalam ke bottleneck performa. Panduan ini memberikan alur kerja menyeluruh (E2E) yang komprehensif untuk menjalankan dan memprofilkan beban kerja JAX di GPU (seperti NVIDIA L4) menggunakan Diagnostik ML Google Cloud dan XProf. Dengan memanfaatkan alat ini, Anda dapat mengidentifikasi operasi yang tidak efisien, mengoptimalkan penggunaan resource komputasi, dan mempercepat proses pelatihan.

Dengan mengikuti panduan ini, Anda akan mempelajari cara:

  1. Instrumentasikan loop pelatihan JAX sederhana untuk pembuatan profil.
  2. Masukkan workload ke dalam container dengan dukungan CUDA yang sesuai.
  3. Deploy workload di Google Kubernetes Engine (GKE) menggunakan JobSet.
  4. Mengambil dan memvisualisasikan profil performa secara dinamis.

Prasyarat

Sebelum memulai, pastikan Anda memiliki:

  • Project Google Cloud yang mengaktifkan penagihan.
  • Cluster GKE dengan dukungan GPU (misalnya, NVIDIA L4).
  • Bucket Google Cloud Storage (GCS) untuk menyimpan profil.
  • CLI gcloud dan kubectl telah diinstal dan dikonfigurasi.
  • Workload Identity dikonfigurasi untuk cluster GKE Anda guna mengakses GCS.

Langkah 1: Mengukur Beban Kerja JAX

Pertama, buat skrip pelatihan JAX (misalnya, train.py). Kita menggunakan SDK google-cloud-mldiagnostics untuk berinteraksi dengan infrastruktur pembuatan profil terkelola.

[!WARNING] Skrip di bawah mencakup loop tak terbatas untuk membuat GPU tetap sibuk selama demonstrasi pembuatan profil sesuai permintaan. Jangan lupa untuk menghentikan tugas secara manual atau menghapus resource GKE setelah Anda selesai untuk menghindari biaya penagihan yang tidak perlu.

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

Langkah 2: Containerisasi (Dockerfile)

Buat Dockerfile untuk mengemas skrip JAX Anda dengan dependensi CUDA yang diperlukan dan ML Diagnostics SDK.

# 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"]

Bangun dan kirim image ke 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

Langkah 3: Deployment (Manifes Kubernetes)

Deploy workload menggunakan JobSet GKE atau Job standar. Untuk mengaktifkan platform Diagnostik ML agar dapat menyisipkan metadata dan merutekan permintaan profil, terapkan label managed-mldiagnostics-gke: "true". Untuk mengetahui detail selengkapnya tentang cara mengonfigurasi GKE untuk Diagnostik ML, lihat Panduan Penyiapan GKE resmi.

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

Terapkan manifes:

kubectl apply -f deploy.yaml

Langkah 4: Pengambilan & Visualisasi

Perekaman Terprogram

Jika Anda menyertakan prof.start() / prof.stop() dalam skrip, profil tersebut akan otomatis diupload ke bucket GCS Anda di jalur: gs://<your-gcs-bucket>/<run-name>/plugins/profile/<session-id>/

Perekaman Sesuai Permintaan

Karena on_demand_xprof=True ditetapkan di machinelearning_run, Anda dapat merekam profil secara dinamis saat tugas sedang berjalan.

Untuk petunjuk mendetail tentang cara menggunakan UI TensorBoard untuk memicu profil sesuai permintaan, memilih pod tertentu, dan melihat rekaman aktivitas yang diambil, lihat dokumentasi publik resmi: Diagnostik ML Google Cloud - Pengambilan profil sesuai permintaan.

Anda juga dapat merekam profil menggunakan gcloud CLI seperti yang dijelaskan dalam Panduan ML Diagnostics CLI.

Dokumentasi publik ini berlaku untuk workload TPU dan GPU yang dikelola oleh ML Diagnostics.