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