پروفایل و اشکال‌زدایی JAX روی GPU با XProf و Nsight Systems

۱. مقدمه

مسیر یادگیری جکس روی پردازنده گرافیکی. آزمایشگاه ۳: پروفایلینگ و اشکال‌زدایی جکس روی پردازنده گرافیکی.

در آزمایشگاه کد «کنترل کامپایل JAX با jax.jit»، یاد گرفتید که کامپایل یکی از دلایلی است که یک برنامه JAX کند به نظر می‌رسد، و چگونه می‌توان از jax.jit برای جلوگیری از توقف کامپایل مجدد استفاده کرد. اما کامپایل تنها یکی از گزینه‌ها است.

وقتی حجم کاری GPU کند است، کد پایتون تقریباً هرگز به شما نمی‌گوید که علت واقعی آن چیست. این می‌تواند کامپایل شدن، اندازه‌گیری زمان‌بندی که هرگز منتظر GPU نمانده، انتقال‌های دستگاه میزبان پنهان در گزارش‌گیری شما، یک دسته داده بسیار کوچک برای مشغول نگه داشتن GPU یا فشار حافظه باشد.

در این آزمایشگاه کد، شما شروع به اندازه‌گیری این جنبه‌ها می‌کنید. شما یک مرحله آموزشی واقعی را پروفایل می‌کنید، ردپا را در XProf می‌خوانید و با Nsight Systems به جدول زمانی CUDA می‌افتید.

کاری که انجام خواهید داد

  • زمان کامپایل اولین فراخوانی را از زمان اجرای کش شده جدا کنید و زمان JAX را با استفاده از block_until_ready() به طور صادقانه تنظیم کنید.
  • ضبط یک مسیر JAX profiler با jax.profiler.trace و حاشیه‌نویسی‌های مرحله‌ای نامگذاری‌شده
  • آن مسیر را در XProf از طریق پورت-فوروارد Cloud Shell و با TensorBoard به عنوان رابط کاربری جایگزین باز کنید.
  • تشخیص چهار الگوی رایج کندی: تعداد زیاد عملیات کوچک، انتقال به دستگاه میزبان، دسته‌های کوچک و فشار بر حافظه
  • با استفاده از Nsight Systems ، یک جدول زمانی در سطح CUDA ضبط کنید و از محدوده‌های NVTX برای علامت‌گذاری ناحیه مورد نظر خود استفاده کنید.
  • گزارش Nsight را درون دفترچه یادداشت با nsys stats خلاصه کنید

آنچه نیاز دارید

  • یک پروژه Google Cloud با قابلیت پرداخت صورتحساب، و اعتبارات کارگاه یا رزرو شامل استفاده از GPU
  • سهمیه حداقل ۲ پردازنده گرافیکی NVIDIA L4 در منطقه انتخابی شما ( نحوه بررسی سهمیه پردازنده گرافیکی )
  • تکمیل آزمایشگاه‌های کد ۱ و ۲ یا یک محیط معادل JAX GPU
  • در صورت تمایل، Nsight Systems روی دستگاه شما نصب می‌شود تا بتوانید گزارش CUDA را در رابط کاربری گرافیکی باز کنید. دانلود آن رایگان است و codelab خلاصه‌ای متنی از همان گزارش را نیز چاپ می‌کند.

زمان تخمینی برای تکمیل: ۷۰ دقیقه .

گردش کار پروفایلینگ

یک راه مفید برای اشکال‌زدایی عملکرد JAX، حرکت از بررسی‌های ساده به سمت ابزارهای عمیق‌تر است.

ابتدا با مسدود کردن تا پایان کار GPU، زمان‌بندی خود را انجام دهید. سپس یک ردیابی JAX profiler ثبت کنید تا بتوانید کامپایل، فعالیت میزبان و اجرای دستگاه را با هم مشاهده کنید. سپس، آن ردیابی را در XProf یا TensorBoard باز کنید تا جدول‌های زمانی، حافظه، نمودارها و آمار عملیات را بررسی کنید. هنگامی که به نمای سطح CUDA نیاز دارید، از Nsight Systems برای دیدن جریان‌ها، هسته‌ها، کپی‌های حافظه، فراخوانی‌های کتابخانه و ارتباطات استفاده کنید.

چه چیزی را جستجو کنیم

وقتی یک مسیر را باز می‌کنید، سعی نکنید هر رویداد را یکجا درک کنید. با بررسی اجمالی چند الگوی بصری رایج شروع کنید.

بازه‌های کامپایل به شما می‌گویند که آیا زمان به جای اجرا، صرف کامپایل XLA می‌شود یا خیر. فواصل خالی در ردیف‌های GPU اغلب به این معنی است که میزبان به اندازه کافی سریع دستگاه را تغذیه نمی‌کند. فعالیت انتقال می‌تواند به همگام‌سازی تصادفی دستگاه میزبان، مانند ثبت وقایع با float(loss) اشاره داشته باشد. اوج‌های حافظه به شما کمک می‌کنند دسته‌ها، فعال‌سازی‌ها یا بافرهای موقت را که GPU را به مرز محدودیت خود نزدیک می‌کنند، تشخیص دهید.

۲. قبل از شروع

پروژه خود را انتخاب کنید

در کنسول گوگل کلود ، یک پروژه با قابلیت پرداخت فعال انتخاب یا ایجاد کنید.

پوسته ابری را باز کنید

برای شروع یک جلسه Cloud Shell ، روی Activate Cloud Shell (آیکون ترمینال در سمت راست بالای کنسول) کلیک کنید، سپس آن را به پروژه خود هدایت کنید:

gcloud config set project <YOUR_PROJECT_ID>

این codelab در همان محیط Codelab 1 اجرا می‌شود: اولین برنامه JAX خود را روی GPUهای NVIDIA با GKE اجرا کنید . اگر کلاستر GKE و JupyterLab Pod شما هنوز در حال اجرا هستند، از مرحله نصب موارد مورد نیاز این codelab صرف نظر کنید. در غیر این صورت، همین حالا محیط را آماده کنید.

فراهم کردن محیط GPU

دستور زیر را در Cloud Shell اجرا کنید.

gcloud services enable \
  container.googleapis.com \
  compute.googleapis.com \
  iam.googleapis.com \
  cloudresourcemanager.googleapis.com \
  logging.googleapis.com \
  monitoring.googleapis.com

git clone https://github.com/Google-Cloud-AI/partner-ai-nvidia.git
cd partner-ai-nvidia/05-workshops/jax-on-gpu/terraform

cp terraform.tfvars.example terraform.tfvars

فایل terraform.tfvars را ویرایش کنید و project_id = " " قرار دهید. project_id = " " ، سپس کلاستر را آماده کرده و JupyterLab را مستقر کنید:

terraform init
terraform apply
$(terraform output -raw get_credentials_command)

cd ..
kubectl apply -f deploy/jupyter.yaml

terraform apply حدود ۱۰ دقیقه طول می‌کشد. پس از اتمام، منتظر Pod و LoadBalancer باشید، سپس توکن یکبار مصرف JupyterLab را از لاگ Pod بخوانید:

kubectl get pod jax-jupyter -w         # wait for Running, then Ctrl+C
kubectl get svc jax-jupyter-svc -w     # wait for EXTERNAL-IP, then Ctrl+C
kubectl logs jax-jupyter | grep -o 'token=[a-z0-9]*' | head -1

http:// :8884 را باز کنید http:// :8884 وارد کنید، توکن را جایگذاری کنید و یک دفترچه یادداشت پایتون ۳ جدید در /workspace ایجاد کنید. هر بلوک کد در این codelab به سلولی از آن دفترچه یادداشت می‌رود.

آنچه این codelab نیاز دارد را نصب کنید

!pip install --quiet xprof nvtx

nsys ، ابزار پروفایلر خط فرمان Nsight Systems، از قبل در محفظه NVIDIA JAX موجود است، بنابراین چیزی برای نصب در جدول زمانی CUDA بعداً وجود ندارد. اگر ترجیح می‌دهید از تب پروفایل TensorBoard به جای XProf مستقل به عنوان نمایشگر ردیابی خود استفاده کنید، tensorboard به آن دستور نصب اضافه کنید.

GPU را تنظیم و تأیید کنید

ابزارها را وارد کنید، GPU را بررسی کنید و کمک‌کننده‌های نمایش کوچک را که بقیه این آزمایشگاه کد از آنها استفاده می‌کند، تعریف کنید. اگر این سلول با مشکل مواجه شد، قبل از ادامه، محیط را اصلاح کنید؛ نمایه‌سازی یک جایگزین CPU گمراه‌کننده است.

import csv
import importlib.util
import io
import os
import pathlib
import shutil
import socket
import subprocess
import sys
import tempfile
import textwrap
import time

from IPython.display import HTML, Javascript, display

import jax
import jax.numpy as jnp
import numpy as np


def require_executable(name):
    """Look up `name` on PATH and assert it's found; returns the absolute path or fails fast."""
    path = shutil.which(name)
    assert path, f"Required executable '{name}' not found on PATH."
    return path


XPROF_BIN = require_executable("xprof")
NSYS_BIN = require_executable("nsys")
TENSORBOARD_BIN = shutil.which("tensorboard")
assert importlib.util.find_spec("nvtx"), "Required Python package 'nvtx' is missing. Install with: pip install nvtx"


def show_bars(rows, title, unit="", lower_is_better=False):
    """Render (label, value) pairs as a horizontal bar chart in HTML, scaled to the largest value."""
    max_value = max(float(value) for _, value in rows) or 1.0
    html = ["<div style='font-family: Arial, sans-serif; max-width: 760px;'>"]
    html.append(f"<h4 style='margin: 0 0 8px 0;'>{title}</h4>")
    for label, value in rows:
        width = max(3, 100 * float(value) / max_value)
        html.append(
            "<div style='display:grid; grid-template-columns: 190px 1fr 115px; gap: 8px; "
            "align-items:center; margin: 6px 0;'>"
            f"<div style='font-size:13px;'>{label}</div>"
            "<div style='background:#f6f8fa; border-radius:6px; overflow:hidden; height:22px;'>"
            f"<div style='height:22px; width:{width:.1f}%; background:#0969da;'></div></div>"
            f"<div style='font-size:13px; font-variant-numeric: tabular-nums;'>{value:.3f} {unit}</div>"
            "</div>"
        )
    html.append(f"<div style='font-size:12px; color:#57606a;'>{'Lower' if lower_is_better else 'Higher'} is better.</div></div>")
    display(HTML("".join(html)))


def show_table(headers, rows, title=None, aligns=None):
    """Render rows as an HTML table; `aligns` is an optional per-column list of "left"/"right"/"center"."""
    aligns = aligns or ["left"] * len(headers)
    html = ["<div style='font-family: system-ui; max-width: 980px;'>"]
    if title:
        html.append(f"<h4 style='margin-bottom: 8px;'>{title}</h4>")
    html.append("<table style='border-collapse: collapse; width: 100%; font-size: 13px;'>")
    html.append("<thead><tr>")
    for h, a in zip(headers, aligns):
        html.append(f"<th style='text-align:{a}; border-bottom:1px solid #d0d7de; padding:6px;'>{h}</th>")
    html.append("</tr></thead><tbody>")
    for row in rows:
        html.append("<tr>")
        for cell, a in zip(row, aligns):
            html.append(f"<td style='text-align:{a}; border-bottom:1px solid #edf0f2; padding:6px; vertical-align:top; white-space:nowrap;'>{cell}</td>")
        html.append("</tr>")
    html.append("</tbody></table></div>")
    display(HTML("".join(html)))

def show_file_list(root, limit=12):
    """Print files under `root` with their sizes (KB); truncates after `limit` entries."""
    root = pathlib.Path(root)
    files = [p for p in sorted(root.rglob("*")) if p.is_file()]
    for p in files[:limit]:
        print(f"{p.stat().st_size / 1024:8.1f} KB  {p.relative_to(root)}")
    if len(files) > limit:
        print(f"... {len(files) - limit} more files")


devices = jax.devices()
gpu_devices = [d for d in devices if d.platform == "gpu"]

print(f"JAX version:     {jax.__version__}")
print(f"Default backend: {jax.default_backend()}")
print(f"Devices:         {devices}")
print(f"xprof:           {XPROF_BIN}")
print(f"tensorboard: {TENSORBOARD_BIN or 'not found, XProf standalone is enough'}")
print(f"nsys:            {NSYS_BIN}")

assert gpu_devices, f"This lesson assumes a GPU backend. Available devices: {devices}"
print(f"GPU devices:     {gpu_devices}")

این سلول مسیرهای تعیین‌شده برای xprof و nsys و لیست دستگاه‌های JAX را چاپ می‌کند، سپس ادعا می‌کند که حداقل یکی از آنها GPU است. اگر جستجوی xprof ناموفق بود، دستور pip install بالا را دوباره اجرا کرده و هسته را مجدداً راه‌اندازی کنید.

۳. یک حجم کاری برای پروفایل ایجاد کنید

شما به چیزی نیاز دارید که ارزش پروفایل کردن داشته باشد، اما به اندازه کافی کوچک باشد که بتوان آن را در عرض چند ثانیه دوباره اجرا کرد. سلول زیر یک MLP دو لایه، یک خطای میانگین مربعات زیان، گرادیان‌ها از jax.value_and_grad و یک به‌روزرسانی SGD ساده را تعریف می‌کند که همه در یک train_step با jitted قرار گرفته‌اند.

چند خط آخر به اندازه مدل اهمیت دارند. فراخوانی warmup یک بار train_step کامپایل می‌کند، بنابراین هر سلول زمان‌بندی که پس از آن می‌آید، به جای اندازه‌گیری تصادفی راه‌اندازی، اجرا را اندازه‌گیری می‌کند.

BATCH = 256
IN_DIM = 1024
HIDDEN = 1024
OUT_DIM = 256
LR = 1e-3


def init_params(key):
    """Initialize the two-layer MLP weights this lesson profiles."""
    k1, k2 = jax.random.split(key)
    return {
        "w1": jax.random.normal(k1, (IN_DIM, HIDDEN), dtype=jnp.float32) * 0.02,
        "w2": jax.random.normal(k2, (HIDDEN, OUT_DIM), dtype=jnp.float32) * 0.02,
    }


def make_batch(key, batch_size=BATCH):
    """Generate a random (x, target) batch with the standard input/output dims."""
    kx, ky = jax.random.split(key)
    x = jax.random.normal(kx, (batch_size, IN_DIM), dtype=jnp.float32)
    y = jax.random.normal(ky, (batch_size, OUT_DIM), dtype=jnp.float32)
    return x, y


def loss_fn(params, batch):
    """Forward pass plus MSE loss"""
    x, target = batch
    hidden = jax.nn.gelu(x @ params["w1"])
    pred = hidden @ params["w2"]
    return jnp.mean((pred - target) ** 2)


# This is the function we will profile.
@jax.jit
def train_step(params, batch):
    """One compiled SGD step"""
    loss, grads = jax.value_and_grad(loss_fn)(params, batch)
    params = jax.tree.map(lambda p, g: p - LR * g, params, grads)
    return params, loss


key = jax.random.key(0)
params = init_params(key)
batch = make_batch(jax.random.fold_in(key, 1))

# Warm up once so later timing is on execution
params, loss = train_step(params, batch)
jax.block_until_ready((params, loss))

print(f"x shape/device:      {batch[0].shape} on {batch[0].device}")
print(f"target shape/device: {batch[1].shape} on {batch[1].device}")
print(f"w1 shape/device:     {params['w1'].shape} on {params['w1'].device}")
print(f"warmup loss:         {float(loss):.4f}")

هر خط از خروجی چاپی باید یک دستگاه CUDA را نامگذاری کند: ورودی‌ها در (256, 1024) ، اهداف در (256, 256) و w1 در (1024, 1024) و به دنبال آن یک مقدار اتلاف گرم شدن.

۴. کامپایل را از اجرا و زمان را صادقانه جدا کنید

یک تابع jitted دو حالت بسیار متفاوت دارد:

  • اولین فراخوانی برای امضای ورودی جدید، ردیابی و کامپایل می‌شود، سپس اجرا می‌گردد.
  • فراخوانی‌های بعدی با همان شکل‌ها و dtypeها، فایل اجرایی کامپایل‌شده را مجدداً استفاده می‌کنند.

تغییر شکل دسته‌ای، امضای جدیدی ایجاد می‌کند، بنابراین JAX مجبور است دوباره کامپایل شود. این دلیل اصلی مکث مداوم بسیاری از حلقه‌های آموزشی است.

هزینه کامپایل را در برابر اجرای ذخیره شده اندازه گیری کنید

در اینجا شما سه بار train_step یکسانی دارید: یک بار با کش پاک شده، یک بار با کش گرم و یک بار با دسته‌ای با شکل متفاوت.

jax.clear_caches()

compile_params = init_params(jax.random.key(101))
compile_batch = make_batch(jax.random.key(102), batch_size=BATCH)

# Trace + compile + execute.
t0 = time.perf_counter()
compile_params, compile_loss = train_step(compile_params, compile_batch)
jax.block_until_ready((compile_params, compile_loss))
first_ms = (time.perf_counter() - t0) * 1000

# Execute using the cached compiled executable.
t0 = time.perf_counter()
compile_params, compile_loss = train_step(compile_params, compile_batch)
jax.block_until_ready((compile_params, compile_loss))
cached_ms = (time.perf_counter() - t0) * 1000

# Different batch shape which triggers another compile.
smaller_batch = make_batch(jax.random.key(103), batch_size=BATCH // 2)
t0 = time.perf_counter()
_, shape_loss = train_step(compile_params, smaller_batch)
shape_change_ms = (time.perf_counter() - t0) * 1000

print(f"First call, same shape (compile + execute): {first_ms:8.2f} ms")
print(f"Cached call, same shape (execute only):     {cached_ms:8.2f} ms")
print(f"New batch shape (compile + execute):        {shape_change_ms:8.2f} ms")
show_bars(
    [
        ("first call", first_ms),
        ("cached call", cached_ms),
        ("new shape", shape_change_ms),
    ],
    title="Compilation Cost vs. Cached Execution",
    unit="ms",
    lower_is_better=True,
)

# Re-warm the original shape.
params, loss = train_step(params, batch)
jax.block_until_ready((params, loss))

شما باید سه عدد و یک نمودار میله‌ای دریافت کنید. دو تا از سه میله - فراخوانی اول و شکل جدید - شامل کامپایل هستند، بنابراین هر دو باید بسیار بزرگتر از فراخوانی ذخیره شده باشند. آن میله میانی تنها میله‌ای است که نشان می‌دهد یک مرحله آموزش حالت پایدار واقعاً چقدر هزینه دارد.

قبل از اینکه ساعت را متوقف کنید، مسدود کنید

JAX کار GPU را به صورت ناهمزمان ارسال می‌کند، بنابراین پایتون می‌تواند قبل از اتمام کار GPU، آن را برگرداند. بدون block_until_ready() ، معمولاً مدت زمانی که پایتون صرف کرده تا کار را در صف قرار دهد را اندازه‌گیری می‌کنید، نه مدت زمانی که GPU برای اجرای آن صرف کرده است.

در اینجا شما دو بار مرحله آموزش یکسانی را دارید: یک بار به اشتباه، و یک بار با یک بلوک روی نتیجه.

t0 = time.perf_counter()
params_dispatch, loss_dispatch = train_step(params, batch)
dispatch_ms = (time.perf_counter() - t0) * 1000

t0 = time.perf_counter()
params_ready, loss_ready = train_step(params, batch)
jax.block_until_ready((params_ready, loss_ready))
ready_ms = (time.perf_counter() - t0) * 1000

print(f"Dispatch-only timing: {dispatch_ms:.3f} ms")
print(f"Blocked timing:       {ready_ms:.3f} ms")
show_bars(
    [("dispatch only", dispatch_ms), ("block_until_ready", ready_ms)],
    title="Timing the Same JAX Step",
    unit="ms",
    lower_is_better=True,
)

نوار فقط اعزام باید کوتاه‌تر از بین این دو باشد. آن عدد کوتاه‌تر، برنامه سریع‌تری نیست، بلکه یک برنامه اندازه‌گیری نشده است.

۵. الگوی «عملیات‌های کوچکِ بیش از حد زیاد» را اصلاح کنید

یکی دیگر از چالش‌های رایج عملکرد GPU، اجرای عملیات کوچک متعدد از پایتون است. هر عملیات کوچک JAX سربار توزیع پایتون را به همراه دارد و ممکن است یک هسته GPU کوچک تولید کند که قبل از اینکه دستگاه به درستی مشغول شود، به پایان می‌رسد.

jax.jit مفید است زیرا به XLA اجازه می‌دهد کل زنجیره را ببیند و آن را به عنوان یک واحد کامپایل شده ترکیب یا زمان‌بندی کند. دو تابع زیر محاسبه یکسانی انجام می‌دهند. اولی عملیات را از پایتون ارسال می‌کند و دومی زنجیره را کامپایل می‌کند.

SMALL_OP_STEPS = 50
small_x = jnp.ones((4096,), dtype=jnp.float32)


def many_small_ops(x):
    """Un-jitted chain of small operations."""
    # Each loop iteration dispatches JAX operations from Python.
    y = x
    for _ in range(SMALL_OP_STEPS):
        y = jnp.sin(y) + 0.01 * y
    return y


@jax.jit
def compiled_chain(x):
    """Same operations under `jax.jit` so XLA can fuse."""
    # JAX traces the whole chain, and XLA can optimize it as one compiled computation.
    y = x
    for _ in range(SMALL_OP_STEPS):
        y = jnp.sin(y) + 0.01 * y
    return y


_ = compiled_chain(small_x).block_until_ready()

# Time the non-jitted chain.
t0 = time.perf_counter()
_ = many_small_ops(small_x).block_until_ready()
small_ops_ms = (time.perf_counter() - t0) * 1000

# Time the compiled chain.
t0 = time.perf_counter()
_ = compiled_chain(small_x).block_until_ready()
compiled_chain_ms = (time.perf_counter() - t0) * 1000

print(f"Many small Python-dispatched ops: {small_ops_ms:8.3f} ms")
print(f"Compiled chain with jax.jit:      {compiled_chain_ms:8.3f} ms")
show_bars(
    [("many small ops", small_ops_ms), ("compiled chain", compiled_chain_ms)],
    title="Too Many Small Operations",
    unit="ms",
    lower_is_better=True,
)

زنجیره کامپایل شده باید میله کوتاه‌تر باشد. آرایه اینجا فقط ۴۰۹۶ عنصر دارد و حلقه ۵۰ بار اجرا می‌شود، بنابراین نسخه بدون jitted، سربار توزیع پایتون را ۵۰ برابر بیشتر برای محاسبات بسیار کم در هر بار پرداخت می‌کند. این الگویی است که باید در کد خودتان تشخیص دهید: یک حلقه پایتون روی عملیات کوچک JAX که می‌تواند به جای آن یک تابع کامپایل شده باشد.

۶. ثبت ردپای JAX profiler

زمان‌بندی‌ها به شما می‌گویند که چیزی کند است. یک ردپا به شما می‌گوید که در طول زمان چه اتفاقی افتاده است : چه زمانی پایتون فعال بوده، چه زمانی XLA کامپایل شده، چه زمانی پردازنده گرافیکی هسته‌ها را اجرا کرده و شکاف‌ها کجا بوده‌اند.

ردپاهای خام پر از نام عملیات سطح پایین هستند، بنابراین قبل از ثبت حجم کار، آن را حاشیه‌نویسی کنید. علامت‌گذاری هر مرحله آموزشی، نشانه‌های قابل خواندن توسط انسان را برای پیمایش در اختیار شما قرار می‌دهد. JAX سه حاشیه‌نویسی ارائه می‌دهد:

  • StepTraceAnnotation مراحل تکرار شده، مانند تکرارهای آموزش، را نامگذاری می‌کند.
  • TraceAnnotation ناحیه‌ای را درون یک مرحله، مانند آماده‌سازی دسته‌ای یا به‌روزرسانی بهینه‌ساز، نامگذاری می‌کند.
  • annotate_function یک تابع پایتون را در trace نامگذاری می‌کند.

این سلول شش مرحله آموزشی حاشیه‌نویسی شده را در یک دایرکتوری موقت ردیابی می‌کند:

trace_dir = pathlib.Path(tempfile.mkdtemp(prefix="jax-trace-"))
trace_params = params
trace_key = key

with jax.profiler.trace(str(trace_dir), create_perfetto_link=False):
    for step_num in range(6):
        with jax.profiler.StepTraceAnnotation("train_step", step_num=step_num):
            with jax.profiler.TraceAnnotation("make_batch"):
                trace_key = jax.random.fold_in(trace_key, step_num + 10)
                trace_batch = make_batch(trace_key)
            with jax.profiler.TraceAnnotation("sgd_update"):
                trace_params, trace_loss = train_step(trace_params, trace_batch)
            jax.block_until_ready((trace_params, trace_loss))

print(f"Trace directory: {trace_dir}")
show_file_list(trace_dir)

شما باید مسیر دایرکتوری ردیابی و به دنبال آن لیست کوتاهی از فایل‌های ردیابی نوشته شده در داخل آن، به همراه اندازه هر کدام را ببینید. این مسیر دایرکتوری همان چیزی است که در مرحله بعدی بیننده را به آن هدایت می‌کنید، بنابراین خروجی سلول را قابل مشاهده نگه دارید.

۷. ردیابی را در XProf باز کنید

XProf رابط کاربری پروفایلر در پشت فرمت ردیابی JAX است. این رابط مستقیماً فهرست ردیابی را می‌خواند و یک رابط وب با نمایشگر ردیابی، نمایشگر حافظه، نمایشگر نمودار و آمار عملیات ارائه می‌دهد.

آن را درون پاد شروع کنید:

xprof_port = 6007
xprof_proc = subprocess.Popen(
    [XPROF_BIN, "--port", str(xprof_port), str(trace_dir)],
    stdout=subprocess.DEVNULL,
    stderr=subprocess.DEVNULL,
)
time.sleep(3)
xprof_url = f"http://localhost:{xprof_port}"
print(f"XProf is running at {xprof_url} (PID {xprof_proc.pid})")
print(f"Stop it later with: kill {xprof_proc.pid}")
display(Javascript(f'window.open("{xprof_url}", "_blank");'))
display(HTML(f'<a href="{xprof_url}" target="_blank" rel="noopener">Open XProf</a>'))

این سلول، آدرس اینترنتی (URL) که XProf روی آن کار می‌کند، شناسه‌ی فرآیند (Process ID) و دستور kill که آن را متوقف می‌کند را چاپ می‌کند. آن PID را یادداشت کنید - در مرحله‌ی پاکسازی (Clean up) به آن نیاز دارید.

از طریق لپ‌تاپ خود با بیننده ارتباط برقرار کنید

به طور پیش‌فرض هیچ چیزی که درون Pod در حال گوش دادن است از دستگاه شما قابل دسترسی نیست. kubectl port-forward یک تونل از Cloud Shell به پورتی درون Pod باز می‌کند و پیش‌نمایش وب Cloud Shell آن تونل را در مرورگر شما منتشر می‌کند.

یک تب جدید Cloud Shell باز کنید - این دستور تا زمانی که آن را متوقف نکنید، در پیش‌زمینه اجرا می‌شود و در این حین، نوت‌بوک شما در پاد به اجرا ادامه می‌دهد:

kubectl port-forward pod/jax-jupyter 8080:6007

سپس در نوار ابزار Cloud Shell روی پیش‌نمایش وب - پیش‌نمایش روی پورت ۸۰۸۰ کلیک کنید تا بتوانید به رابط کاربری دسترسی پیدا کنید.

ردپا را بخوانید

به محض اینکه تب پیش‌نمایش وب، XProf را بارگذاری کرد:

  1. از منوی کشویی Runs، Run مورد نظر را انتخاب کنید.
  2. ابزارها را باز کنید - trace_viewer .
  3. محدوده‌های train_step را پیدا کنید - اینها نام‌های StepTraceAnnotation از مرحله قبل هستند.
  4. به دنبال بازه‌های کامپایل، شکاف‌های GPU، کپی‌ها و فعالیت هسته باشید.

چک لیست خواندن ردیابی

جدول زیر نشان می‌دهد که چگونه الگوی بصری در نمودار، اقدام بعدی را آغاز می‌کند:

سرنخ بصری

احتمالاً به چه معناست

چه چیزی را بعداً امتحان کنیم

طول کامپایل طولانی قبل از مرحله اول

کامپایل JIT با فراخوانی اول معمولی

قبل از اندازه گیری گرم کنید.

کامپایل کردن span های بین مراحل

کامپایل مجدد

تغییر شکل‌ها، dtypeها یا آرگومان‌های استاتیک را از codelab 2 بررسی کنید.

ردیف‌های GPU دارای شکاف‌های سفید هستند

میزبان به پردازنده گرافیکی (GPU) تغذیه نمی‌دهد

به دنبال بارگذاری داده‌ها، چاپ، float(loss) و تبدیل NumPy باشید.

بسیاری از هسته‌های ریز

راه‌اندازی سربار یا مناطق کامپایل شده خیلی کوچک

JIT یک تابع بزرگتر؛ کار دسته‌ای.

فعالیت Memcpy بین مراحل

انتقال‌های میزبان-دستگاه

معیارها را روی دستگاه نگه دارید؛ کمتر وارد سیستم شوید؛ از محافظ انتقال استفاده کنید.

حافظه با اوج بالا

ممکن است فعال‌سازی‌ها یا بافرهای موقت غالب باشند

memory_viewer باز کنید و دسته کوچکتری را امتحان کنید.

سه نکته ورودی دیگر نیز ارزش ذکر دارند، هرچند این آزمایشگاه کد از آنها استفاده نمی‌کند: start_trace() و stop_trace() برای نواحی ردیابی برنامه‌نویسی که در بلوک with jax.profiler.trace(...) قرار نمی‌گیرند، start_server() به همراه python -m jax.collect_profile برای پروفایل کردن jobهای طولانی مدت، و Perfetto export برای ردیابی‌های باز در رابط کاربری Perfetto. راهنمای پروفایل کردن JAX هر سه مورد را پوشش می‌دهد.

۸. تشخیص انتقال میزبان و اندازه دسته

دو مورد از موارد چک لیست بالا آنقدر رایج هستند که ارزش تکرار دارند: انتقال تصادفی میزبان، و دسته‌ای که برای پر کردن GPU خیلی کوچک است.

انتقال‌های میزبان-دستگاه

برگرداندن یک مقدار JAX به پایتون درون یک حلقه، همگام‌سازی را اجباری می‌کند. نمونه‌های رایج آن float(loss) ، .item() ، np.asarray(...) و چاپ آرایه‌ها هستند.

این کد دو سبک گزارش‌گیری را مقایسه می‌کند. نسخه بد، زیان را در هر مرحله به یک float پایتون تبدیل می‌کند. نسخه بهتر، زیان‌ها را به صورت آرایه‌های JAX نگه می‌دارد و در انتها یک بار همگام‌سازی می‌کند.

N = 30

# Convert the loss to a Python float every step.
t0 = time.perf_counter()
bad_params = params
bad_losses = []
for _ in range(N):
    bad_params, bad_loss = train_step(bad_params, batch)
    bad_losses.append(float(bad_loss))
bad_ms = (time.perf_counter() - t0) * 1000 / N

# Keep metrics as JAX arrays and synchronize once.
t0 = time.perf_counter()
good_params = params
good_losses = []
for _ in range(N):
    good_params, good_loss = train_step(good_params, batch)
    good_losses.append(good_loss)
jax.block_until_ready((good_params, good_losses))
good_ms = (time.perf_counter() - t0) * 1000 / N

print(f"float(loss) every step: {bad_ms:.3f} ms / step")
print(f"deferred sync:          {good_ms:.3f} ms / step")
show_bars(
    [("float(loss) every step", bad_ms), ("deferred sync", good_ms)],
    title="Cost of Pulling Metrics to Python",
    unit="ms/step",
    lower_is_better=True,
)

# Transfer guard can catch accidental transfers.
try:
    _, guard_loss = train_step(params, batch)
    guard_loss.block_until_ready()
    with jax.transfer_guard("disallow"):
        _ = float(guard_loss)
except RuntimeError as e:
    print("\nTransfer guard caught an implicit transfer:")
    print(str(e).splitlines()[0])

هزینه هر مرحله برای نسخه float(loss) باید از بین دو میله، بالاترین باشد و سلول باید با چاپ اولین خط RuntimeError که محافظ انتقال ایجاد کرده است، کار خود را تمام کند.

اندازه دسته و توان عملیاتی

دسته‌های کوچک اغلب به پردازنده گرافیکی (GPU) کار موازی کافی نمی‌دهند. معمولاً با افزایش دسته، توان عملیاتی بهبود می‌یابد و سپس به محض اشباع پردازنده گرافیکی یا محدودیت حافظه، روند آن مسطح می‌شود.

اندازه دسته نیز بر حافظه تأثیر می‌گذارد. پیمایش زیر شامل یک تخمین ساده برای بافرهای دسته‌ای شکل در این گذر رو به جلو است: ورودی، فعال‌سازی پنهان و خروجی. آموزش واقعی از چیزی بیش از این استفاده می‌کند، زیرا گرادیان‌ها و بافرهای موقت نیز حساب می‌شوند.

@jax.jit
def forward_only(params, x):
    """Forward pass only."""
    hidden = jax.nn.gelu(x @ params["w1"])
    return hidden @ params["w2"]


batch_results = []
bytes_per_float32 = np.dtype(np.float32).itemsize

for batch_size in (1, 8, 32, 128, 256, 512, 1024):
    xb = jax.random.normal(jax.random.key(batch_size), (batch_size, IN_DIM), dtype=jnp.float32)
    _ = forward_only(params, xb).block_until_ready()

    reps = 80 if batch_size <= 128 else 30
    # Async dispatch lets JAX queue all `reps` calls without blocking. We block once at the end and divide by reps, so this measures *amortized* time per call when the executor stays busy.
    t0 = time.perf_counter()
    for _ in range(reps):
        yb = forward_only(params, xb)
    yb.block_until_ready()

    ms = (time.perf_counter() - t0) * 1000 / reps
    examples_per_sec = batch_size / (ms / 1000)

    # Simple memory estimate for batch-shaped forward buffers.
    estimated_forward_mib = batch_size * (IN_DIM + HIDDEN + OUT_DIM) * bytes_per_float32 / 2**20
    batch_results.append((batch_size, ms, examples_per_sec, estimated_forward_mib))

show_table(
    ["Batch", "ms/call (avg)", "examples/sec", "estimated batch buffers"],
    [
        (batch_size, f"{ms:.3f}", f"{examples_per_sec:,.0f}", f"{estimated_mib:.1f} MiB")
        for batch_size, ms, examples_per_sec, estimated_mib in batch_results
    ],
    title="Batch Size, Throughput, and Memory",
    aligns=["right", "right", "right", "right"],
)

show_bars(
    [(f"batch {batch_size}", examples_per_sec) for batch_size, _, examples_per_sec, _ in batch_results],
    title="Throughput by Batch Size",
    unit="ex/s",
    lower_is_better=False,
)

شما یک جدول با هفت ردیف و یک نمودار میله‌ای دریافت می‌کنید. به ستون examples/sec نگاه کنید، نه ms/call : توان عملیاتی باید در چند اندازه دسته اول به شدت افزایش یابد و سپس با اشباع GPU، مسطح شود. ستون ms/call در کل مسیر رشد می‌کند، که انتظار می‌رود.

۹. فشار حافظه GPU را بررسی کنید

JAX معمولاً در اولین استفاده، درصد مشخصی از حافظه GPU را از قبل اختصاص می‌دهد. این کار عمدی است زیرا سربار تخصیص و قطعه قطعه شدن را کاهش می‌دهد. همچنین به این معنی است که nvidia-smi حتی زمانی که مدل شما کوچک است، می‌تواند تقریباً پر به نظر برسد، که باعث می‌شود افراد زیادی به دنبال نشت حافظه‌ای باشند که وجود ندارد.

برای مشاهده سریع سطح فرآیند memory_stats و برای تجزیه و تحلیل عمیق‌تر از ابزارهای حافظه XProf استفاده کنید.

def gib(value):
    """Convert a byte count to gibibytes."""
    return value / 2**30


print(f"{'device':<26} {'limit':>12} {'in use':>12} {'peak':>12}")
print("-" * 66)
for device in gpu_devices:
    stats = device.memory_stats()
    if not stats:
        print(f"{str(device):<26} memory_stats unavailable")
        continue
    limit = stats.get("bytes_limit")
    in_use = stats.get("bytes_in_use")
    peak = stats.get("peak_bytes_in_use")
    limit_s = f"{gib(limit):.2f} GiB" if limit is not None else "n/a"
    in_use_s = f"{gib(in_use):.2f} GiB" if in_use is not None else "n/a"
    peak_s = f"{gib(peak):.2f} GiB" if peak is not None else "n/a"
    print(f"{str(device):<26} {limit_s:>12} {in_use_s:>12} {peak_s:>12}")

شما باید به ازای هر پردازنده گرافیکی قابل مشاهده، یک ردیف داشته باشید. ستون limit ، رزرو تخصیص‌دهنده است نه کل حافظه کارت، و in use برای این MLP کوچک باید بخش کوچکی از آن باشد.

تنظیمات حافظه باید قبل از وارد کردن JAX تنظیم شوند. در دفترچه یادداشت، این به معنی تنظیم آنها قبل از راه‌اندازی هسته و سپس راه‌اندازی مجدد هسته است.

متغیر

مثال

استفاده از زمانی که

XLA_PYTHON_CLIENT_MEM_FRACTION

0.50

شما یک پردازنده گرافیکی (GPU) را به اشتراک می‌گذارید و می‌خواهید JAX حافظه کمتری را رزرو کند.

XLA_PYTHON_CLIENT_PREALLOCATE

false

شما تخصیص بر اساس تقاضا را می‌خواهید و ریسک تکه‌تکه شدن بیشتری را می‌پذیرید.

XLA_PYTHON_CLIENT_ALLOCATOR

platform

شما در حال اشکال‌زدایی حافظه هستید و می‌خواهید تخصیص حافظه را آزاد کنید؛ برای آموزش عادی خیلی کند است.

این سلول نشان می‌دهد که کدام یک از آنها در هسته فعلی تنظیم شده‌اند:

for name in (
    "XLA_PYTHON_CLIENT_MEM_FRACTION",
    "XLA_PYTHON_CLIENT_PREALLOCATE",
    "XLA_PYTHON_CLIENT_ALLOCATOR",
):
    print(f"{name}={os.environ.get(name, '<unset>')}")

print("\nExample for a shared GPU, set before launching Python/Jupyter:")
print("export XLA_PYTHON_CLIENT_MEM_FRACTION=0.50")

به احتمال زیاد هر سه چاپ خواهند شد ، به این معنی که شما روی پیش‌فرض‌ها هستید: پیش‌تخصیص روی، کسر ۷۵٪.

۱۰. ضبط جدول زمانی CUDA با Nsight Systems

XProf اولین پروفایلر مناسب برای JAX است. Nsight Systems دومین view است، برای زمانی که به جدول زمانی CUDA نیاز دارید: استریم‌ها، فراخوانی‌های API CUDA، هسته‌ها، نسخه‌های حافظه، فراخوانی‌های cuBLAS و cuDNN و در نهایت NCCL.

گردش کار چهار بخش دارد:

  1. حجم کار را در یک متن کوتاه بگنجانید.
  2. محدوده‌های NVTX را در اطراف مراحلی که برایتان مهم هستند اضافه کنید.
  3. اسکریپت را با nsys profile اجرا کنید.
  4. فایل .nsys-rep را با رابط کاربری گرافیکی Nsight Systems باز کنید.

اسکریپت ضبط را بنویسید

پروفایل کردن مستقیم یک نوت‌بوک آسان نیست، بنابراین کد زیر به جای آن یک اسکریپت مستقل کوچک می‌نویسد. این اسکریپت یک بار فعال می‌شود، سپس از محدوده‌های NVTX برای علامت‌گذاری ناحیه‌ای که ارزش ثبت شدن دارد استفاده می‌کند.

لازم نیست تک تک خطوط را بخوانید. بیشتر آن همان مدل کوچک قبلی را تکرار می‌کند تا nsys بتواند یک فرآیند پایتون جدید را نمایه‌سازی کند.

nsight_dir = pathlib.Path(tempfile.mkdtemp(prefix="jax-nsight-"))
script_path = nsight_dir / "nsight_train_step.py"
report_base = nsight_dir / "jax_train_step"
report_path = pathlib.Path(f"{report_base}.nsys-rep")

script_source = f"""
import nvtx
import jax
import jax.numpy as jnp

BATCH = {BATCH}
IN_DIM = {IN_DIM}
HIDDEN = {HIDDEN}
OUT_DIM = {OUT_DIM}
LR = {LR}


def init_params(key):
    k1, k2 = jax.random.split(key)
    return {{ "{{" }}
        "w1": jax.random.normal(k1, (IN_DIM, HIDDEN), dtype=jnp.float32) * 0.02,
        "w2": jax.random.normal(k2, (HIDDEN, OUT_DIM), dtype=jnp.float32) * 0.02,
    {{ "}}" }}


def make_batch(key):
    kx, ky = jax.random.split(key)
    return (
        jax.random.normal(kx, (BATCH, IN_DIM), dtype=jnp.float32),
        jax.random.normal(ky, (BATCH, OUT_DIM), dtype=jnp.float32),
    )


def loss_fn(params, batch):
    x, target = batch
    hidden = jax.nn.gelu(x @ params["w1"])
    pred = hidden @ params["w2"]
    return jnp.mean((pred - target) ** 2)


@jax.jit
def train_step(params, batch):
    loss, grads = jax.value_and_grad(loss_fn)(params, batch)
    params = jax.tree.map(lambda p, g: p - LR * g, params, grads)
    return params, loss


key = jax.random.key(0)
params = init_params(key)
batch = make_batch(key)

# Warm up before the region we care about.
params, loss = train_step(params, batch)
jax.block_until_ready((params, loss))

with nvtx.annotate("profile_region", domain="jax_course"):
    for step in range(8):
        with nvtx.annotate(f"train_step_{{ "{{" }}step{{ "}}" }}", domain="jax_course"):
            key = jax.random.fold_in(key, step)
            batch = make_batch(key)
            params, loss = train_step(params, batch)
            jax.block_until_ready((params, loss))
"""
script_path.write_text(textwrap.dedent(script_source))

print(f"Wrote script: {script_path}")
print(f"Report path:  {report_path}")

سلول دو مسیری را که استفاده خواهد کرد چاپ می‌کند. هنوز هیچ چیزی اجرا نشده است.

سیستم‌های Nsight را اجرا کنید

فراخوانی زیر فقط زمانی شروع به جمع‌آوری می‌کند که محدوده NVTX profile_region@jax_course شروع شود و زمانی که آن محدوده به پایان برسد متوقف می‌شود. این کار باعث می‌شود گزارش به جای دفن شدن در زیر بار راه‌اندازی فرآیند و کامپایل JIT، بر اجرای پس از گرم شدن متمرکز بماند.

import base64

nsys_cmd = [
    NSYS_BIN,
    "profile",
    "--trace=cuda,nvtx,osrt,cudnn,cublas",
    "--capture-range=nvtx",
    "--capture-range-end=stop",
    "--nvtx-capture=profile_region@jax_course",
    "--force-overwrite=true",
    f"--output={report_base}",
    "-e",
    "NSYS_NVTX_PROFILER_REGISTER_ONLY=0",
    sys.executable,
    str(script_path),
]

print("Running Nsight Systems:")
print(" ".join(nsys_cmd))
nsys_result = subprocess.run(nsys_cmd, capture_output=True, text=True, check=False)

if nsys_result.returncode != 0:
    print("nsys profile failed")
    print(nsys_result.stdout[-2000:])
    print(nsys_result.stderr[-4000:])
    raise RuntimeError("Nsight Systems capture failed. This lesson requires nsys to run successfully.")

print(f"Nsight report: {report_path}  ({report_path.stat().st_size / 1024:.1f} KB)")

b64 = base64.b64encode(report_path.read_bytes()).decode()
download_html = (
    f'<a href="data:application/octet-stream;base64,{b64}" '
    f'download="{report_path.name}" '
    f'style="font-weight:600;">Download {report_path.name}</a>'
)

display(HTML(download_html))

ضبط کمی طول می‌کشد، زیرا یک فرآیند پایتون جدید را اجرا می‌کند که JAX را وارد کرده و train_step دوباره کامپایل می‌کند. وقتی کار تمام شد، مسیر و اندازه گزارش و یک لینک دانلود دریافت می‌کنید. آن لینک یک URL داده است، بنابراین مستقیماً از نوت‌بوک در مرورگر شما کار می‌کند - نیازی به پورت فوروارد نیست.

جدول زمانی Nsight را بخوانید

برای باز کردن گزارش در رابط کاربری گرافیکی:

  1. Nsight Systems را از طریق developer.nvidia.com/nsight-systems روی دستگاه محلی خود نصب کنید. این برنامه رایگان است و روی لینوکس، macOS و ویندوز اجرا می‌شود.
  2. گزارش .nsys-rep تولید شده را با استفاده از لینک بالا دانلود کنید.
  3. nsys-ui را اجرا کنید و فایل را با استفاده از منوی File - Open باز کنید، یا آن را به داخل پنجره بکشید و رها کنید.

با این ردیف‌ها شروع کنید:

  • NVTX : پیدا کردن profile_region و train_step_* .
  • ردیف‌های پردازنده گرافیکی CUDA : پوشش هسته و شکاف‌های سفید را بررسی کنید، که در آن شکاف سفید به معنای زمان بیکاری پردازنده گرافیکی است.
  • API CUDA : به دنبال فراخوانی‌های cudaLaunchKernel ، cudaMemcpy* و همگام‌سازی باشید.
  • cuBLAS / cuDNN : فراخوانی‌های matmul، convolution و attention library هنگام ردیابی در اینجا ظاهر می‌شوند.
  • زمان اجرای سیستم عامل ( osrt ) : انتظار میزبان، قفل شدن، خواب و سایر رفتارهای مسدودکننده.

برای معیارهای استفاده از GPU، اگر محیط شما امکان جمع‌آوری معیارهای GPU را فراهم می‌کند، ضبط را با --gpu-metrics-devices=cuda-visible دوباره اجرا کنید.

خلاصه کردن گزارش

رابط کاربری گرافیکی (GUI) تجربه اصلی Nsight است، اما nsys stats خلاصه متنی مفیدی ارائه می‌دهد که می‌توانید آن را بخوانید. این کد، جداول اصلی هسته، عملیات حافظه و CUDA API را چاپ می‌کند.

stats_cmd = [
    NSYS_BIN,
    "stats",
    "--quiet",
    "--report",
    "cuda_gpu_kern_sum,cuda_gpu_mem_time_sum,cuda_api_sum",
    "--format",
    "csv",
    "--timeunit",
    "msec",
    str(report_path),
]

stats_result = subprocess.run(stats_cmd, capture_output=True, text=True, check=False)
if stats_result.returncode != 0:
    print("nsys stats failed")
    print(stats_result.stderr[-3000:])
else:
    sections = [s for s in stats_result.stdout.strip().split("\n\n") if s.strip()]
    for idx, section in enumerate(sections[:3], start=1):
        reader = csv.DictReader(io.StringIO(section))
        rows = list(reader)
        if not rows:
            continue
        headers = reader.fieldnames or []
        print(f"\nReport section {idx}: {headers}")
        compact_rows = []
        for row in rows[:8]:
            name = row.get("Name") or row.get("Operation") or row.get("Name:Demangled") or ""
            time_pct = row.get("Time (%)", "")
            total = row.get("Total Time (ms)") or row.get("Total Time (msec)") or row.get("Total Time") or row.get("Total Time (ns)") or ""
            count = row.get("Instances") or row.get("Count") or row.get("Num Calls") or row.get("Calls") or ""
            compact_rows.append((time_pct, total, count, name[:90]))
        show_table(["Time (%)", "Total time (ms)", "Count", "Name / operation"], compact_rows, title=f"Nsight stats section {idx}", aligns=["right", "right", "right", "left"])

شما باید تا سه جدول داشته باشید که هر کدام هشت ردیف برتر را بر اساس زمان نشان می‌دهند. ابتدا خلاصه هسته را بخوانید: این جدول هسته‌های GPU را بر اساس کل زمان رتبه‌بندی می‌کند، که به شما می‌گوید دستگاه واقعاً آن را کجا صرف کرده است. آن را با جدول عملیات حافظه مقایسه کنید - اگر کپی‌ها با هسته‌ها بر سر زمان رقابت کنند، به مشکل انتقال میزبان از قبل برمی‌گردید.

۱۱. تمیز کردن

گره GPU چه چیزی روی آن اجرا کنید و چه نکنید، هزینه‌ها را محاسبه می‌کند، بنابراین این مرحله را نادیده نگیرید.

بینندگانی را که ابتدا در داخل پاد شروع کرده‌اید، متوقف کنید: اجرای kill دستوراتی که سلول‌های XProf و TensorBoard چاپ می‌کنند، و در هر تب Cloud Shell که هنوز kubectl port-forward در حال اجرا است، Ctrl+C را فشار دهید.

حجم کاری Jupyter، شامل LoadBalancer و درایو دائمی را حذف کنید:

kubectl delete -f deploy/jupyter.yaml

کلاستر، Node Pool، VPC و حساب کاربری سرویس را از بین ببرید:

cd terraform
terraform destroy

وقتی از شما خواسته شد، yes را تایپ کنید، سپس تأیید کنید که چیزی جا نمانده است:

gcloud container clusters list
gcloud compute instances list

هر دو باید برای این پروژه خالی باشند. اگر فقط برای این مجموعه پروژه‌ای ایجاد کرده‌اید، می‌توانید کل پروژه را از کنسول Cloud حذف کنید.

۱۲. تبریک

شما یک مرحله آموزشی JAX را از ابتدا تا انتها، از اندازه‌گیری time.perf_counter تا جدول زمانی CUDA، نمایه‌سازی کردید.

آنچه آموخته‌اید

  • چگونه می‌توان JAX را با استفاده از block_until_ready() صادقانه زمان‌بندی کرد و چرا اندازه‌گیری صرفاً مبتنی بر اعزام بی‌معنی است؟
  • چگونه زمان کامپایل اولین فراخوانی را از زمان اجرای ذخیره شده در حافظه پنهان جدا کنیم، و چگونه یک شکل ورودی جدید، کامپایل مجدد را الزامی می‌کند
  • نحوه ثبت یک ردیابی با jax.profiler.trace و برچسب‌گذاری آن با StepTraceAnnotation و TraceAnnotation
  • چگونه می‌توان آن ردیابی را در XProf و در TensorBoard به عنوان یک رابط کاربری جایگزین، از طریق پورت-فوروارد Cloud Shell باز کرد؟
  • چگونه الگوهای رایج کندی را تشخیص دهیم: کامپایل مجدد، تعداد زیاد عملیات کوچک، انتقال به دستگاه میزبان، اندازه‌های دسته‌ای ناکارآمد و فشار حافظه
  • نحوه بررسی حافظه GPU با memory_stats() و نحوه تغییر رزرو JAX XLA_PYTHON_CLIENT_MEM_FRACTION قبل از شروع به کار
  • نحوه ضبط و خواندن جدول زمانی CUDA با Nsight Systems با استفاده از محدوده‌های NVTX، و نحوه خلاصه کردن آن با nsys stats

مراحل بعدی

  • Codelab 4: آموزش یک مدل روی GPU با JAX، Optax و Fashion-MNIST که در آن یک حلقه آموزشی واقعی را روی داده‌های واقعی تجربه خواهید کرد، با استفاده از عادت‌های پروفایل‌سازی از این آزمایشگاه که از قبل وجود دارند.
  • ضبط Nsight را با اضافه کردن --gpu-metrics-devices=cuda-visible به nsys_cmd دوباره اجرا کنید و شمارنده‌های استفاده از GPU را با پوشش هسته که در جدول زمانی مشاهده کردید مقایسه کنید.
  • قبل از راه‌اندازی مجدد هسته، مقدار XLA_PYTHON_CLIENT_MEM_FRACTION=0.50 را تنظیم کنید، سپس سلول memory_stats() را دوباره اجرا کنید و بررسی کنید که ستون limit چگونه تغییر کرده است.

اسناد مرجع