XProf と ML Diagnostics を使用して GPU で JAX をプロファイリングする

GPU で大規模な JAX モデルを最適化するには、パフォーマンスのボトルネックを詳細に把握する必要があります。このガイドでは、Google Cloud ML Diagnostics と 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 CLI と kubectl CLI がインストールされ、構成されている。
  • GCS にアクセスするように構成された GKE クラスタの Workload Identity。

ステップ 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 に push します。

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: Deployment(Kubernetes マニフェスト)

GKE JobSet または標準の Job を使用してワークロードをデプロイします。ML Diagnostics プラットフォームがメタデータを挿入してプロファイル リクエストを転送できるようにするには、ラベル 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() を含めた場合、これらのプロファイルはパス gs://<your-gcs-bucket>/<run-name>/plugins/profile/<session-id>/ の下にある GCS バケットに自動的にアップロードされます。

オンデマンド キャプチャ

on_demand_xprof=True は machinelearning_run で設定されているため、ジョブの実行中にプロファイルを動的にキャプチャできます。

TensorBoard UI を使用してオンデマンド プロファイルをトリガーし、特定の Pod を選択して、キャプチャされたトレースを表示する詳しい手順については、公式の公開ドキュメント(Google Cloud ML Diagnostics - オンデマンド プロファイルのキャプチャ)をご覧ください。

ML Diagnostics CLI ガイドで説明されているように、gcloud CLI を使用してプロファイルをキャプチャすることもできます。

この公開ドキュメントは、ML Diagnostics によって管理される TPU ワークロードと GPU ワークロードの両方に適用されます。