जीपीयू पर बड़े पैमाने पर JAX मॉडल को ऑप्टिमाइज़ करने के लिए, परफ़ॉर्मेंस से जुड़ी समस्याओं के बारे में पूरी जानकारी होना ज़रूरी है. इस गाइड में, Google Cloud ML Diagnostics और XProf का इस्तेमाल करके, जीपीयू (जैसे, NVIDIA L4) पर JAX वर्कलोड चलाने और उनकी प्रोफ़ाइलिंग करने के बारे में पूरी जानकारी दी गई है. इन टूल का इस्तेमाल करके, इन कामों को पूरा किया जा सकता है: बेकार कार्रवाइयों की पहचान करना, कंप्यूट रिसोर्स के इस्तेमाल को ऑप्टिमाइज़ करना, और ट्रेनिंग को तेज़ी से पूरा करना.
इस गाइड को पढ़कर, आपको इन कामों को करने का तरीका पता चलेगा:
- प्रोफ़ाइलिंग के लिए, एक आसान JAX ट्रेनिंग लूप इंस्ट्रुमेंट करें.
- वर्कलोड को सही CUDA सपोर्ट के साथ कंटेनर में रखें.
- JobSet का इस्तेमाल करके, Google Kubernetes Engine (GKE) पर वर्कलोड डिप्लॉय करें.
- परफ़ॉर्मेंस प्रोफ़ाइल को डाइनैमिक तौर पर कैप्चर और विज़ुअलाइज़ करें.
ज़रूरी शर्तें
शुरू करने से पहले, पक्का करें कि आपके पास ये चीज़ें हों:
- बिलिंग की सुविधा वाला Google क्लाउड प्रोजेक्ट.
- जीपीयू की सुविधा वाला GKE क्लस्टर (जैसे, NVIDIA L4).
- प्रोफ़ाइलें सेव करने के लिए, Google Cloud Storage (GCS) बकेट.
gcloudऔरkubectlसीएलआई इंस्टॉल और कॉन्फ़िगर किए गए हों.- GCS को ऐक्सेस करने के लिए, आपके GKE क्लस्टर के लिए Workload Identity कॉन्फ़िगर किया गया हो.
पहला चरण: JAX वर्कलोड को इंस्ट्रुमेंट करना
सबसे पहले, 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()
दूसरा चरण: कंटेनर बनाना (Dockerfile)
ज़रूरी CUDA डिपेंडेंसी और ML Diagnostics SDK के साथ अपनी JAX स्क्रिप्ट को पैकेज करने के लिए, Dockerfile बनाएं.
# 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
तीसरा चरण: डिप्लॉयमेंट (Kubernetes मेनिफ़ेस्ट)
GKE JobSet या स्टैंडर्ड जॉब का इस्तेमाल करके, वर्कलोड को डिप्लॉय करें. एमएल डाइग्नोस्टिक्स प्लैटफ़ॉर्म को मेटाडेटा इंजेक्ट करने और प्रोफ़ाइल के अनुरोधों को रूट करने की अनुमति देने के लिए, managed-mldiagnostics-gke: "true" लेबल लागू करें. एमएल डाइग्नोस्टिक्स के लिए GKE को कॉन्फ़िगर करने के बारे में ज़्यादा जानने के लिए, 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
चौथा चरण: डेटा इकट्ठा करना और उसे विज़ुअलाइज़ करना
प्रोग्राम के हिसाब से, अपने-आप होने वाली प्रोसेस के ज़रिए कैप्चर करना
अगर आपने अपनी स्क्रिप्ट में 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.
gcloud सीएलआई का इस्तेमाल करके भी प्रोफ़ाइलें कैप्चर की जा सकती हैं. इसके बारे में एमएल डाइग्नोस्टिक्स सीएलआई गाइड में बताया गया है.
यह सार्वजनिक दस्तावेज़, ML Diagnostics की मदद से मैनेज किए जाने वाले टीपीयू और जीपीयू, दोनों के वर्कलोड पर लागू होता है.