Lập hồ sơ JAX trên GPU bằng XProf và ML Diagnostics

Việc tối ưu hoá các mô hình JAX quy mô lớn trên GPU đòi hỏi phải có thông tin chi tiết về các nút thắt cổ chai hiệu suất. Hướng dẫn này cung cấp một quy trình toàn diện, từ đầu đến cuối (E2E) để chạy và lập hồ sơ khối lượng công việc JAX trên GPU (chẳng hạn như NVIDIA L4) bằng Google Cloud ML Diagnostics và XProf. Bằng cách tận dụng những công cụ này, bạn có thể xác định các hoạt động không hiệu quả, tối ưu hoá việc sử dụng tài nguyên điện toán và tăng tốc các lượt chạy huấn luyện.

Bằng cách làm theo hướng dẫn này, bạn sẽ tìm hiểu cách:

  1. Gắn mã theo dõi một vòng lặp huấn luyện JAX đơn giản để phân tích.
  2. Chứa khối lượng công việc trong vùng chứa với sự hỗ trợ CUDA phù hợp.
  3. Triển khai tải công việc trên Google Kubernetes Engine (GKE) bằng JobSet.
  4. Ghi lại và trực quan hoá các hồ sơ hiệu suất một cách linh động.

Điều kiện tiên quyết

Trước khi bắt đầu, hãy đảm bảo rằng bạn có:

  • Một dự án trên Google Cloud đã bật tính năng thanh toán.
  • Một cụm GKE có hỗ trợ GPU (ví dụ: NVIDIA L4).
  • Một bộ chứa Google Cloud Storage (GCS) để lưu trữ hồ sơ.
  • Đã cài đặt và định cấu hình các CLI gcloud và kubectl.
  • Workload Identity được định cấu hình cho cụm GKE của bạn để truy cập vào GCS.

Bước 1: Đo lường khối lượng công việc JAX

Trước tiên, hãy tạo một tập lệnh huấn luyện JAX (ví dụ: train.py). Chúng ta sẽ sử dụng SDK google-cloud-mldiagnostics để tương tác với cơ sở hạ tầng lập hồ sơ được quản lý.

[!WARNING] Tập lệnh bên dưới có một vòng lặp vô hạn để giữ cho GPU luôn bận trong các bản minh hoạ lập hồ sơ theo yêu cầu. Hãy nhớ dừng công việc theo cách thủ công hoặc xoá các tài nguyên GKE sau khi bạn hoàn tất để tránh các khoản phí không cần thiết.

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()

Bước 2: Tạo vùng chứa (Dockerfile)

Tạo một Dockerfile để đóng gói tập lệnh JAX bằng các phần phụ thuộc CUDA bắt buộc và 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"]

Tạo và chuyển hình ảnh vào 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

Bước 3: Triển khai (Tệp kê khai Kubernetes)

Triển khai khối lượng công việc bằng JobSet hoặc Job tiêu chuẩn của GKE. Để cho phép nền tảng ML Diagnostics chèn siêu dữ liệu và định tuyến các yêu cầu về hồ sơ, hãy áp dụng nhãn managed-mldiagnostics-gke: "true". Để biết thêm thông tin về cách định cấu hình GKE cho ML Diagnostics, hãy tham khảo Hướng dẫn thiết lập GKE chính thức.

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

Áp dụng tệp kê khai:

kubectl apply -f deploy.yaml

Bước 4: Thu thập và trực quan hoá

Chụp ảnh có lập trình

Nếu bạn thêm prof.start() / prof.stop() vào tập lệnh, thì những hồ sơ đó sẽ tự động được tải lên bộ chứa GCS của bạn theo đường dẫn: gs://<your-gcs-bucket>/<run-name>/plugins/profile/<session-id>/

Ghi hình theo yêu cầu

Vì on_demand_xprof=True được đặt trong machinelearning_run, nên bạn có thể ghi lại các hồ sơ một cách linh động trong khi công việc đang chạy.

Để biết hướng dẫn chi tiết về cách sử dụng giao diện người dùng TensorBoard để kích hoạt hồ sơ theo yêu cầu, chọn các nhóm cụ thể và xem các dấu vết đã ghi lại, vui lòng tham khảo tài liệu công khai chính thức: Google Cloud ML Diagnostics – Ghi lại hồ sơ theo yêu cầu.

Bạn cũng có thể ghi lại hồ sơ bằng gcloud CLI như mô tả trong Hướng dẫn về CLI chẩn đoán ML.

Tài liệu công khai này áp dụng cho cả khối lượng công việc TPU và GPU do ML Diagnostics quản lý.