การสร้างโปรไฟล์ JAX บน GPU ด้วย XProf และการวินิจฉัย ML

การเพิ่มประสิทธิภาพโมเดล JAX ขนาดใหญ่ใน GPU ต้องมีการมองเห็นอย่างละเอียดเกี่ยวกับ คอขวดด้านประสิทธิภาพ คู่มือนี้มีเวิร์กโฟลว์แบบครบวงจร (E2E) ที่ครอบคลุมสำหรับการเรียกใช้และการสร้างโปรไฟล์ภาระงาน JAX ใน GPU (เช่น NVIDIA L4) โดยใช้การวินิจฉัย ML ของ Google Cloud และ XProf การใช้เครื่องมือเหล่านี้จะช่วยให้คุณ ระบุการดำเนินการที่ไม่มีประสิทธิภาพ เพิ่มประสิทธิภาพการใช้ทรัพยากรการประมวลผล และเร่ง การเรียกใช้การฝึก

เมื่อทำตามคู่มือนี้ คุณจะได้เรียนรู้วิธีการต่างๆ ต่อไปนี้

  1. วัดคุมลูปการฝึก JAX อย่างง่ายสำหรับการสร้างโปรไฟล์
  2. สร้างคอนเทนเนอร์ให้กับภาระงานด้วยการรองรับ CUDA ที่เหมาะสม
  3. ติดตั้งใช้งานภาระงานใน Google Kubernetes Engine (GKE) โดยใช้ JobSet
  4. บันทึกและแสดงภาพโปรไฟล์ประสิทธิภาพแบบไดนามิก

ข้อกำหนดเบื้องต้น

ก่อนเริ่มต้น โปรดตรวจสอบว่าคุณมีสิ่งต่อไปนี้

  • โปรเจ็กต์ Google Cloud ที่เปิดใช้การเรียกเก็บเงิน
  • คลัสเตอร์ GKE ที่รองรับ GPU (เช่น NVIDIA L4)
  • Bucket ของ Google Cloud Storage (GCS) สำหรับจัดเก็บโปรไฟล์
  • ติดตั้งและกำหนดค่า CLI ของ gcloud และ kubectl แล้ว
  • Workload Identity ที่กำหนดค่าไว้สำหรับคลัสเตอร์ GKE เพื่อเข้าถึง GCS

ขั้นตอนที่ 1: ติดตั้งเครื่องมือเวิร์กโหลด JAX

ก่อนอื่น ให้สร้างสคริปต์การฝึก JAX (เช่น train.py) เราใช้ SDK google-cloud-mldiagnostics เพื่อโต้ตอบกับโครงสร้างพื้นฐานการจัดโปรไฟล์ที่มีการจัดการ

[!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 พร้อมกับทรัพยากร Dependency ของ CUDA ที่จำเป็นและ 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"]

สร้างและพุชอิมเมจไปยัง 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 Manifest)

ติดตั้งใช้งานเวิร์กโหลดโดยใช้ JobSet หรือ Job มาตรฐานของ GKE หากต้องการเปิดใช้แพลตฟอร์ม ML Diagnostics เพื่อแทรกข้อมูลเมตาและกำหนดเส้นทางคำขอโปรไฟล์ ให้ใช้ป้ายกำกับ managed-mldiagnostics-gke: "true" ดูรายละเอียดเพิ่มเติมเกี่ยวกับการกำหนดค่า GKE สำหรับการวินิจฉัย ML ได้ที่คู่มือการตั้งค่า 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

ใช้ไฟล์ Manifest

kubectl apply -f deploy.yaml

ขั้นตอนที่ 4: การบันทึกและการแสดงภาพ

การจับภาพแบบเป็นโปรแกรม

หากคุณใส่ prof.start() / prof.stop() ไว้ในสคริปต์ ระบบจะอัปโหลดโปรไฟล์เหล่านั้นไปยัง Bucket ของ GCS โดยอัตโนมัติภายใต้เส้นทางต่อไปนี้ gs://<your-gcs-bucket>/<run-name>/plugins/profile/<session-id>/

การจับภาพออนดีมานด์

เนื่องจากตั้งค่า on_demand_xprof=True ไว้ใน machinelearning_run คุณจึงบันทึก โปรไฟล์แบบไดนามิกได้ขณะที่งานกำลังทำงาน

ดูวิธีการโดยละเอียดเกี่ยวกับวิธีใช้ UI ของ TensorBoard เพื่อทริกเกอร์โปรไฟล์ตามต้องการ เลือกพ็อดที่เฉพาะเจาะจง และดูการติดตามที่บันทึกไว้ได้ที่ เอกสารประกอบสาธารณะอย่างเป็นทางการ การวินิจฉัย ML ของ Google Cloud - การบันทึกโปรไฟล์ตามต้องการ

นอกจากนี้ คุณยังบันทึกโปรไฟล์โดยใช้ gcloud CLI ได้ตามที่อธิบายไว้ในคู่มือ CLI การวินิจฉัย ML

เอกสารประกอบสาธารณะนี้ใช้ได้กับทั้งเวิร์กโหลด TPU และ GPU ที่จัดการโดย ML Diagnostics