XProf 및 ML 진단을 사용하여 GPU에서 JAX 프로파일링

GPU에서 대규모 JAX 모델을 최적화하려면 성능 병목 현상을 자세히 파악해야 합니다. 이 가이드에서는 Google Cloud ML 진단 및 XProf를 사용하여 GPU (예: NVIDIA L4)에서 JAX 워크로드를 실행하고 프로파일링하는 포괄적인 엔드 투 엔드 (E2E) 워크플로를 제공합니다. 이러한 도구를 활용하면 비효율적인 작업을 식별하고, 컴퓨팅 리소스 사용량을 최적화하고, 학습 실행을 가속화할 수 있습니다.

이 가이드를 따르면 다음 작업을 수행하는 방법을 알 수 있습니다.

  1. 프로파일링을 위해 간단한 JAX 학습 루프를 계측합니다.
  2. 적절한 CUDA 지원으로 워크로드를 컨테이너화합니다.
  3. JobSet을 사용하여 Google Kubernetes Engine (GKE)에 워크로드를 배포합니다.
  4. 성능 프로필을 동적으로 캡처하고 시각화합니다.

기본 요건

시작하기 전에 다음 사항을 확인하세요.

  • 결제가 사용 설정된 Google Cloud 프로젝트.
  • GPU 지원 (예: NVIDIA L4)이 있는 GKE 클러스터
  • 프로필을 저장할 Google Cloud Storage (GCS) 버킷
  • gcloud 및 kubectl CLI가 설치되고 구성되어 있어야 합니다.
  • 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 워크로드 모두에 적용됩니다.