إنشاء ملفات تعريف JAX على وحدات معالجة الرسومات باستخدام XProf وML Diagnostics

يتطلّب تحسين نماذج JAX الكبيرة الحجم على وحدات معالجة الرسومات إمكانية الاطّلاع بشكل مفصّل على المشاكل التي تؤدي إلى بطء الأداء. يقدّم هذا الدليل سير عمل شاملاً ومتكاملاً (E2E) لتشغيل أحمال عمل JAX وإنشاء ملفات تعريف لها على وحدات معالجة الرسومات (مثل NVIDIA L4) باستخدام Google Cloud ML Diagnostics وXProf. من خلال الاستفادة من هذه الأدوات، يمكنك تحديد العمليات غير الفعّالة، وتحسين استخدام موارد الحوسبة، وتسريع عمليات التدريب.

من خلال اتّباع هذا الدليل، ستتعرّف على كيفية:

  1. تجهيز حلقة تدريب بسيطة في JAX لإنشاء ملفات تعريف
  2. تضمين حجم المعالجة في حاوية مع توفير دعم CUDA المناسب
  3. نشر عبء العمل على Google Kubernetes Engine (GKE) باستخدام JobSet
  4. تسجيل ملفات الأداء وعرضها بشكل مرئي بطريقة ديناميكية

المتطلبات الأساسية

قبل البدء، تأكَّد من توفّر ما يلي:

  • مشروع Google Cloud تم تفعيل الفوترة فيه
  • مجموعة GKE متوافقة مع وحدة معالجة الرسومات (مثل NVIDIA L4)
  • حزمة Google Cloud Storage (GCS) لتخزين الملفات الشخصية
  • تم تثبيت واجهتَي سطر الأوامر gcloud وkubectl وضبطهما.
  • تم ضبط ميزة Workload Identity لمجموعة GKE من أجل الوصول إلى GCS.

الخطوة 1: قياس أداء عبء عمل JAX

أولاً، أنشئ برنامجًا نصيًا للتدريب باستخدام JAX (مثل train.py). نستخدم حزمة تطوير البرامج (SDK) google-cloud-mldiagnostics للتفاعل مع البنية الأساسية المُدارة لإنشاء الملفات الشخصية.

[!WARNING] يتضمّن البرنامج النصي أدناه حلقة لا نهائية لإبقاء وحدة معالجة الرسومات مشغولة من أجل عرض توضيحي لعملية إنشاء الملفات الشخصية عند الطلب. تذكَّر إيقاف المهمة يدويًا أو حذف موارد 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 التابعة المطلوبة وحزمة تطوير البرامج (SDK) لأداة &quot;أدوات تشخيص تعلُّم الآلة&quot;.

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

يمكنك نشر عبء العمل باستخدام JobSet في GKE أو Job عادي. لتفعيل منصة ML Diagnostics من أجل إدخال البيانات الوصفية وتوجيه طلبات الملفات الشخصية، طبِّق التصنيف managed-mldiagnostics-gke: "true". لمزيد من التفاصيل حول إعداد GKE لأداة ML Diagnostics، يُرجى الرجوع إلى دليل إعداد 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 ضمن المسار: gs://<your-gcs-bucket>/<run-name>/plugins/profile/<session-id>/

تسجيل البيانات عند الطلب

بما أنّ on_demand_xprof=True تم ضبطه في machinelearning_run، يمكنك تسجيل الملفات الشخصية بشكل ديناميكي أثناء تشغيل المهمة.

للحصول على تعليمات مفصّلة حول كيفية استخدام واجهة مستخدم TensorBoard لتفعيل الملفات الشخصية عند الطلب واختيار وحدات معيّنة وعرض عمليات التتبُّع التي تم تسجيلها، يُرجى الرجوع إلى المستندات الرسمية المتاحة للجميع: أداة Google Cloud ML Diagnostics - تسجيل الملفات الشخصية عند الطلب.

يمكنك أيضًا تسجيل الملفات الشخصية باستخدام gcloud CLI كما هو موضّح في دليل ML Diagnostics CLI.

تنطبق هذه المستندات المتاحة للجميع على أحمال عمل وحدات المعالجة المركزية (TPU) ووحدات معالجة الرسومات (GPU) التي تديرها أداة ML Diagnostics.