GPU에서 대규모 JAX 모델을 최적화하려면 성능 병목 현상을 자세히 파악해야 합니다. 이 가이드에서는 Google Cloud ML 진단 및 XProf를 사용하여 GPU (예: NVIDIA L4)에서 JAX 워크로드를 실행하고 프로파일링하는 포괄적인 엔드 투 엔드 (E2E) 워크플로를 제공합니다. 이러한 도구를 활용하면 비효율적인 작업을 식별하고, 컴퓨팅 리소스 사용량을 최적화하고, 학습 실행을 가속화할 수 있습니다.
이 가이드를 따르면 다음 작업을 수행하는 방법을 알 수 있습니다.
- 프로파일링을 위해 간단한 JAX 학습 루프를 계측합니다.
- 적절한 CUDA 지원으로 워크로드를 컨테이너화합니다.
- JobSet을 사용하여 Google Kubernetes Engine (GKE)에 워크로드를 배포합니다.
- 성능 프로필을 동적으로 캡처하고 시각화합니다.
기본 요건
시작하기 전에 다음 사항을 확인하세요.
- 결제가 사용 설정된 Google Cloud 프로젝트.
- GPU 지원 (예: NVIDIA L4)이 있는 GKE 클러스터
- 프로필을 저장할 Google Cloud Storage (GCS) 버킷
gcloud및kubectlCLI가 설치되고 구성되어 있어야 합니다.- GCS에 액세스하도록 GKE 클러스터에 구성된 워크로드 아이덴티티
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)
필수 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
3단계: 배포 (Kubernetes 매니페스트)
GKE JobSet 또는 표준 작업을 사용하여 워크로드를 배포합니다. ML 진단 플랫폼이 메타데이터를 삽입하고 프로필 요청을 라우팅하도록 하려면 managed-mldiagnostics-gke: "true" 라벨을 적용하세요. ML 진단을 위해 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()를 포함한 경우 이러한 프로필은 gs://<your-gcs-bucket>/<run-name>/plugins/profile/<session-id>/ 경로 아래의 GCS 버킷에 자동으로 업로드됩니다.
온디맨드 캡처
on_demand_xprof=True은 machinelearning_run에 설정되어 있으므로 작업이 실행되는 동안 프로필을 동적으로 캡처할 수 있습니다.
TensorBoard UI를 사용하여 주문형 프로필을 트리거하고, 특정 포드를 선택하고, 캡처된 트레이스를 보는 방법에 관한 자세한 안내는 공식 공개 문서인 Google Cloud ML 진단 - 주문형 프로필 캡처를 참고하세요.
ML 진단 CLI 가이드에 설명된 대로 gcloud CLI를 사용하여 프로필을 캡처할 수도 있습니다.
이 공개 문서는 ML 진단으로 관리되는 TPU 및 GPU 워크로드 모두에 적용됩니다.