如要在 GPU 上最佳化大型 JAX 模型,就必須深入瞭解效能瓶頸。本指南提供完整的端對端 (E2E) 工作流程,說明如何使用 Google Cloud ML Diagnostics 和 XProf,在 GPU (例如 NVIDIA L4) 上執行及剖析 JAX 工作負載。善用這些工具,即可找出效率不彰的作業、充分運用運算資源,並加快訓練執行速度。
本指南將說明如何:
- 檢測簡單的 JAX 訓練迴圈以進行剖析。
- 將工作負載容器化,並提供適當的 CUDA 支援。
- 使用 JobSet 在 Google Kubernetes Engine (GKE) 上部署工作負載。
- 動態擷取及以圖像方式呈現效能剖析檔。
必要條件
開始前,請先確認下列事項:
- 已啟用計費功能的 Google Cloud 專案。
- 支援 GPU 的 GKE 叢集 (例如 NVIDIA L4)。
- 用來儲存設定檔的 Google Cloud Storage (GCS) bucket。
- 已安裝及設定
gcloud和kubectlCLI。 - 為 GKE 叢集設定 Workload Identity,以便存取 GCS。
步驟 1:檢測 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()
步驟 2:容器化 (Dockerfile)
建立 Dockerfile,將 JAX 指令碼與必要的 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 資訊清單)
使用 GKE JobSet 或標準 Job 部署工作負載。如要讓 ML 診斷平台注入中繼資料並轉送剖析要求,請套用標籤 managed-mldiagnostics-gke: "true"。如要進一步瞭解如何為 ML Diagnostics 設定 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
步驟 4:擷取及視覺化
程式輔助擷取
如果指令碼中包含 prof.start() / prof.stop(),這些設定檔會自動上傳至 GCS bucket,路徑如下:gs://<your-gcs-bucket>/<run-name>/plugins/profile/<session-id>/
隨選擷取
由於 on_demand_xprof=True 是在 machinelearning_run 中設定,因此您可以在工作執行期間動態擷取設定檔。
如需如何使用 TensorBoard UI 觸發隨選設定檔、選取特定 Pod,以及查看擷取的追蹤記錄,請參閱官方公開說明文件:Google Cloud ML Diagnostics - On-demand profile capture。
您也可以使用 gcloud CLI 擷取剖析資料,詳情請參閱 ML Diagnostics CLI 指南。
這份公開說明文件適用於由機器學習診斷功能管理的 TPU 和 GPU 工作負載。