GPU'larda büyük ölçekli JAX modellerini optimize etmek için performans darboğazları hakkında ayrıntılı bilgi gerekir. Bu kılavuzda, Google Cloud ML Diagnostics ve XProf kullanarak GPU'larda (ör. NVIDIA L4) JAX iş yüklerini çalıştırma ve profillendirme için kapsamlı bir uçtan uca (E2E) iş akışı sağlanır. Bu araçlardan yararlanarak verimsiz işlemleri belirleyebilir, bilgi işlem kaynağı kullanımını optimize edebilir ve eğitim çalıştırmalarınızı hızlandırabilirsiniz.
Bu kılavuzu takip ederek şunları yapmayı öğreneceksiniz:
- Profil oluşturma için basit bir JAX eğitim döngüsü oluşturun.
- İş yükünü uygun CUDA desteğiyle kapsülleyin.
- İş yükünü JobSet kullanarak Google Kubernetes Engine'e (GKE) dağıtın.
- Performans profillerini dinamik olarak yakalayıp görselleştirin.
Ön koşullar
Başlamadan önce şunlara sahip olduğunuzdan emin olun:
- Faturalandırmanın etkin olduğu bir Google Cloud projesi.
- GPU desteği (ör. NVIDIA L4) olan bir GKE kümesi.
- Profilleri depolamak için bir Google Cloud Storage (GCS) paketi.
gcloudvekubectlCLI'ları yüklenmiş ve yapılandırılmış olmalıdır.- GCS'ye erişmek için GKE kümeniz için Workload Identity yapılandırılmış olmalıdır.
1. adım: JAX iş yükünü izleme
Öncelikle bir JAX eğitim komut dosyası oluşturun (ör. train.py). Yönetilen profilleme altyapısıyla etkileşim kurmak için google-cloud-mldiagnostics SDK'sını kullanırız.
[!UYARI] Aşağıdaki komut dosyası, isteğe bağlı profil oluşturma gösterimleri için GPU'yu meşgul tutmak üzere sonsuz bir döngü içerir. Gereksiz faturalandırma maliyetlerinden kaçınmak için işi manuel olarak durdurmayı veya GKE kaynaklarını tamamladıktan sonra silmeyi unutmayın.
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. adım: Kapsayıcılaştırma (Dockerfile)
JAX komut dosyanızı gerekli CUDA bağımlılıkları ve ML Diagnostics SDK ile paketlemek için Dockerfile oluşturun.
# 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"]
Görüntüyü derleyip Artifact Registry'ye gönderin:
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. adım: Dağıtım (Kubernetes Manifest)
GKE JobSet veya standart bir iş kullanarak iş yükünü dağıtın. Makine öğrenimi teşhis platformunun meta veri eklemesini ve profil isteklerini yönlendirmesini sağlamak için managed-mldiagnostics-gke: "true" etiketini uygulayın. GKE'yi makine öğrenimi teşhisleri için yapılandırma hakkında daha fazla bilgi edinmek üzere resmi GKE Kurulum Kılavuzu'na bakın.
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 dosyasını uygulayın:
kubectl apply -f deploy.yaml
4. adım: Yakalama ve görselleştirme
Programatik Yakalama
Komut dosyanıza prof.start() / prof.stop() eklediyseniz bu profiller
otomatik olarak GCS paketinizdeki şu yola yüklenir:
gs://<your-gcs-bucket>/<run-name>/plugins/profile/<session-id>/
İsteğe bağlı çekim
on_demand_xprof=True, machinelearning_run içinde ayarlandığından iş çalışırken profilleri dinamik olarak yakalayabilirsiniz.
TensorBoard kullanıcı arayüzünü kullanarak isteğe bağlı profilleri tetikleme, belirli pod'ları seçme ve yakalanan izleri görüntüleme hakkında ayrıntılı talimatlar için lütfen resmi herkese açık dokümanları inceleyin: Google Cloud ML Diagnostics - On-demand profile capture.
Ayrıca, ML Diagnostics CLI Kılavuzu'nda açıklandığı gibi gcloud CLI'yı kullanarak da profilleri yakalayabilirsiniz.
Bu herkese açık doküman, ML Diagnostics tarafından yönetilen hem TPU hem de GPU iş yükleri için geçerlidir.