XProf এবং Nsight Systems ব্যবহার করে GPU-তে JAX প্রোফাইল এবং ডিবাগ করুন

১. ভূমিকা

জিপিইউ-তে জ্যাক্স শেখার পথ। ল্যাব ৩: জিপিইউ-তে জ্যাক্স-এর প্রোফাইলিং এবং ডিবাগিং।

"jax.jit দিয়ে JAX কম্পাইলেশন নিয়ন্ত্রণ করুন" কোডল্যাবে আপনি শিখেছেন যে, একটি JAX প্রোগ্রাম ধীরগতির মনে হওয়ার অন্যতম কারণ হলো কম্পাইলেশন, এবং রিকম্পাইল বন্ধ করার জন্য কীভাবে jax.jit ব্যবহার করতে হয়। কিন্তু কম্পাইলেশনই একমাত্র সম্ভাব্য কারণ নয়।

যখন কোনো জিপিইউ ওয়ার্কলোড ধীরগতির হয়, তখন পাইথন কোড প্রায় কখনোই আপনাকে বলে না যে এর আসল কারণটি কী। এর কারণ হতে পারে কম্পাইলেশন, এমন কোনো টাইমিং পরিমাপ যা জিপিইউ-এর জন্য অপেক্ষা করেনি, আপনার লগিং-এর আড়ালে লুকিয়ে থাকা হোস্ট-ডিভাইস ট্রান্সফার, জিপিইউ-কে ব্যস্ত রাখার জন্য ব্যাচটি খুব ছোট হওয়া, অথবা মেমোরির উপর চাপ।

এই কোডল্যাবে আপনি এই দিকগুলো পরিমাপ করা শুরু করবেন। আপনি একটি বাস্তব ট্রেনিং স্টেপের প্রোফাইল তৈরি করবেন, XProf-এ ট্রেসটি পড়বেন এবং Nsight Systems-এর সাহায্যে CUDA টাইমলাইনটি বিস্তারিতভাবে দেখবেন।

আপনি যা করবেন

  • প্রথম-কল কম্পাইলেশন সময়কে ক্যাশড এক্সিকিউশন সময় থেকে আলাদা করুন, এবং block_until_ready() ব্যবহার করে JAX-এর সময় সততার সাথে পরিমাপ করুন।
  • jax.profiler.trace এবং named step annotations ব্যবহার করে একটি JAX profiler ট্রেস ক্যাপচার করুন।
  • বিকল্প ফ্রন্ট এন্ড হিসেবে TensorBoard ব্যবহার করে, একটি ক্লাউড শেল পোর্ট-ফরোয়ার্ডের মাধ্যমে XProf- এ সেই ট্রেসটি খুলুন।
  • ধীরগতির চারটি সাধারণ ধরণ নির্ণয় করুন: অত্যধিক ছোট ছোট অপারেশন, হোস্ট-ডিভাইস ট্রান্সফার, ছোট ব্যাচ এবং মেমরি প্রেসার।
  • Nsight Systems ব্যবহার করে একটি CUDA-লেভেল টাইমলাইন ক্যাপচার করুন, এবং আপনার প্রয়োজনীয় অঞ্চলটি চিহ্নিত করতে NVTX রেঞ্জ ব্যবহার করুন।
  • নোটবুকের ভিতরে nsys stats ব্যবহার করে Nsight রিপোর্টটির সারসংক্ষেপ করুন।

আপনার যা যা লাগবে

  • বিলিং সক্ষম একটি গুগল ক্লাউড প্রজেক্ট, এবং ওয়ার্কশপ ক্রেডিট অথবা জিপিইউ ব্যবহারের জন্য একটি রিজার্ভেশন।
  • আপনার নির্বাচিত অঞ্চলে কমপক্ষে ২টি এনভিডিয়া এল৪ জিপিইউ-এর জন্য কোটা ( জিপিইউ কোটা কীভাবে চেক করবেন )
  • কোডল্যাব ১ ও ২ সম্পন্ন করা, অথবা একটি সমতুল্য JAX GPU পরিবেশ
  • ঐচ্ছিকভাবে, আপনার নিজের মেশিনে Nsight Systems ইনস্টল করা যেতে পারে, যাতে আপনি GUI-তে CUDA রিপোর্টটি খুলতে পারেন। এটি বিনামূল্যে ডাউনলোড করা যায় এবং কোডল্যাবটি একই রিপোর্টের একটি টেক্সট সারাংশও প্রিন্ট করে।

সম্পূর্ণ করতে আনুমানিক সময়: ৭০ মিনিট

প্রোফাইলিং ওয়ার্কফ্লো

JAX পারফরম্যান্স ডিবাগিং করার একটি কার্যকর উপায় হলো সাধারণ পরীক্ষা থেকে আরও গভীর টুলের দিকে অগ্রসর হওয়া।

প্রথমে GPU-এর কাজ শেষ না হওয়া পর্যন্ত ব্লক করে আপনার টাইমিং ঠিক করুন। তারপর একটি JAX প্রোফাইলার ট্রেস ক্যাপচার করুন, যাতে আপনি কম্পাইলেশন, হোস্ট অ্যাক্টিভিটি এবং ডিভাইস এক্সিকিউশন একসাথে দেখতে পারেন। এরপর, টাইমলাইন, মেমরি, গ্রাফ এবং অপারেশন স্ট্যাটিস্টিকস পরীক্ষা করার জন্য সেই ট্রেসটি XProf বা TensorBoard-এ খুলুন। যখন আপনার CUDA-স্তরের ভিউ প্রয়োজন হবে, তখন স্ট্রিম, কার্নেল, মেমরি কপি, লাইব্রেরি কল এবং কমিউনিকেশন দেখার জন্য Nsight Systems ব্যবহার করুন।

কী খুঁজতে হবে

যখন আপনি একটি ট্রেস খুলবেন, তখন একবারে প্রতিটি ইভেন্ট বোঝার চেষ্টা করবেন না। কয়েকটি সাধারণ ভিজ্যুয়াল প্যাটার্ন খুঁজে বের করার মাধ্যমে শুরু করুন।

কম্পাইল স্প্যান আপনাকে বলে দেয় যে এক্সিকিউশনের পরিবর্তে XLA কম্পাইলেশনে সময় ব্যয় হচ্ছে কিনা। GPU রো-তে খালি ফাঁকা স্থান প্রায়শই বোঝায় যে হোস্ট ডিভাইসটিকে যথেষ্ট দ্রুত ডেটা সরবরাহ করছে না। ট্রান্সফার অ্যাক্টিভিটি অনিচ্ছাকৃত হোস্ট-ডিভাইস সিনক্রোনাইজেশনের দিকে ইঙ্গিত করতে পারে, যেমন float(loss) দিয়ে লগিং করা। মেমোরি পিক আপনাকে ব্যাচ, অ্যাক্টিভেশন বা টেম্পোরারি বাফার শনাক্ত করতে সাহায্য করে, যা GPU-কে তার ক্ষমতার প্রায় শেষ সীমায় ঠেলে দিচ্ছে।

২. শুরু করার আগে

আপনার প্রকল্প নির্বাচন করুন

গুগল ক্লাউড কনসোলে , বিলিং সক্ষম করা আছে এমন একটি প্রজেক্ট নির্বাচন করুন বা তৈরি করুন।

ওপেন ক্লাউড শেল

একটি ক্লাউড শেল সেশন শুরু করতে অ্যাক্টিভেট ক্লাউড শেল (কনসোলের উপরের ডানদিকে থাকা টার্মিনাল আইকন)-এ ক্লিক করুন, তারপর এটিকে আপনার প্রজেক্টে নির্দেশ করুন:

gcloud config set project <YOUR_PROJECT_ID>

এই কোডল্যাবটি কোডল্যাব ১: Run your first JAX program on NVIDIA GPUs with GKE-এর মতো একই পরিবেশে চলে। যদি আপনার GKE ক্লাস্টার এবং JupyterLab Pod এখনও চালু থাকে, তাহলে সরাসরি "Install what this codelab needs" অংশে চলে যান। অন্যথায়, এখনই পরিবেশটি প্রস্তুত করুন।

জিপিইউ পরিবেশ প্রস্তুত করুন

ক্লাউড শেলে নিম্নলিখিতটি চালান।

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 = " " তারপর project_id = " " প্রস্তুত করুন এবং JupyterLab স্থাপন করুন:

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

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

terraform apply প্রায় ১০ মিনিট সময় লাগে। এটি শেষ হলে, Pod এবং LoadBalancer-এর জন্য অপেক্ষা করুন, তারপর Pod log থেকে ওয়ান-টাইম JupyterLab টোকেনটি পড়ুন:

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 এ একটি নতুন পাইথন ৩ নোটবুক তৈরি করুন। এই কোডল্যাবের প্রতিটি কোড ব্লক সেই নোটবুকের একটি সেলে যাবে।

এই কোডল্যাবের জন্য যা যা প্রয়োজন তা ইনস্টল করুন।

!pip install --quiet xprof nvtx

nsys , অর্থাৎ Nsight Systems-এর কমান্ড-লাইন প্রোফাইলার, আগে থেকেই NVIDIA JAX কন্টেইনারের সাথে অন্তর্ভুক্ত থাকে, তাই পরবর্তীতে CUDA টাইমলাইনের জন্য এটি ইনস্টল করার কিছু নেই। যদি আপনি আপনার ট্রেস ভিউয়ার হিসেবে স্বতন্ত্র XProf-এর পরিবর্তে TensorBoard প্রোফাইল ট্যাব ব্যবহার করতে চান, তাহলে ওই ইনস্টল কমান্ডের সাথে tensorboard যোগ করুন।

জিপিইউ সেট আপ এবং যাচাই করুন

টুলগুলো ইম্পোর্ট করুন, জিপিইউ (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), একটি মিন-স্কোয়ার্ড-এরর লস (mean-squared-error loss), jax.value_and_grad থেকে প্রাপ্ত গ্রেডিয়েন্ট এবং একটি সাধারণ এসজিডি (SGD) আপডেট সংজ্ঞায়িত করা হয়েছে, যা সবই একটি জিটেড train_step jitted train_step)-এর মধ্যে মোড়ানো আছে।

মডেলের মতোই শেষের কয়েকটি লাইনও সমান গুরুত্বপূর্ণ। ওয়ার্মআপ কলটি 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) অবস্থানে থাকা টার্গেটগুলো, এবং (1024, 1024) অবস্থানে থাকা w1 , এবং এর পরে একটিমাত্র ওয়ার্মআপ লস ভ্যালু থাকতে হবে।

৪. কম্পাইলেশনকে এক্সিকিউশন থেকে আলাদা করুন এবং সততার সাথে সময় নির্ধারণ করুন।

একটি জিটেড ফাংশনের দুটি সম্পূর্ণ ভিন্ন মোড রয়েছে:

  • নতুন ইনপুট সিগনেচারের জন্য প্রথম অনুরোধটি ট্রেস ও কম্পাইল করার পর এক্সিকিউট হয়।
  • পরবর্তীতে একই আকার ও ডেটাটাইপ ব্যবহার করে করা কলগুলো কম্পাইল করা এক্সিকিউটেবলটি পুনরায় ব্যবহার করে।

ব্যাচ শেপ পরিবর্তন করলে একটি নতুন সিগনেচার তৈরি হয়, ফলে 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,
)

শুধুমাত্র প্রেরণের বারটি দুটির মধ্যে ছোটটি হওয়া উচিত। এই ছোট সংখ্যাটি কোনো দ্রুততর প্রোগ্রাম নয়, এটি একটি অপরিমাপযোগ্য সংখ্যা।

৫. অতিরিক্ত ছোট ছোট অপারেশনের প্যাটার্নটি ঠিক করুন

জিপিইউ পারফরম্যান্সের আরেকটি সাধারণ চ্যালেঞ্জ হলো পাইথন থেকে অনেকগুলো ছোট ছোট অপারেশন চালু করা। প্রতিটি ছোট JAX অপারেশনে পাইথন ডিসপ্যাচ ওভারহেড থাকে এবং এর ফলে একটি ছোট জিপিইউ কার্নেল তৈরি হতে পারে, যা ডিভাইসটি সঠিকভাবে ব্যস্ত হওয়ার আগেই শেষ হয়ে যায়।

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

কম্পাইল করা চেইনটি ছোট বার হওয়া উচিত। এখানকার অ্যারেটিতে মাত্র ৪০৯৬টি এলিমেন্ট আছে এবং লুপটি ৫০ বার চলে, তাই আন-জিটেড ভার্সনটি প্রতিবার খুব সামান্য গাণিতিক কাজের জন্য পাইথনের ডিসপ্যাচ ওভারহেড ৫০ বার বহন করে। আপনার নিজের কোডে এই প্যাটার্নটিই চিনতে হবে: ছোট ছোট JAX অপারেশনের জন্য একটি পাইথন লুপ, যেগুলো একটি কম্পাইল করা ফাংশন দিয়েই করা যেত।

৬. একটি JAX প্রোফাইলার ট্রেস ক্যাপচার করুন

টাইমিং আপনাকে বলে দেয় যে কোনো কিছু ধীরগতিতে চলছে। একটি ট্রেস আপনাকে জানায় সময়ের সাথে সাথে কী ঘটেছে : কখন পাইথন সক্রিয় ছিল, কখন XLA কম্পাইল হয়েছে, কখন GPU কার্নেল চালিয়েছে এবং কোথায় কোথায় বিরতি ছিল।

র ট্রেসগুলো নিম্ন-স্তরের অপারেশনের নামে পরিপূর্ণ থাকে, তাই ওয়ার্কলোড ক্যাপচার করার আগে সেটিকে অ্যানোটেট করুন। প্রতিটি ট্রেনিং স্টেপ চিহ্নিত করলে আপনি নেভিগেট করার জন্য সহজে পাঠযোগ্য নির্দেশক চিহ্ন পাবেন। JAX তিনটি অ্যানোটেশন অফার করে:

  • StepTraceAnnotation পুনরাবৃত্ত ধাপগুলোর নাম দেয়, যেমন প্রশিক্ষণের পুনরাবৃত্তি।
  • TraceAnnotation একটি স্টেপের ভেতরের কোনো অঞ্চলের নাম দেয়, যেমন ব্যাচ প্রিপ বা অপটিমাইজার আপডেট।
  • annotate_function ট্রেস-এ একটি পাইথন ফাংশনের নাম উল্লেখ করে।

এই সেলটি ছয়টি টীকাযুক্ত প্রশিক্ষণ ধাপকে একটি অস্থায়ী ডিরেক্টরিতে চিহ্নিত করে:

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 হলো JAX ট্রেস ফরম্যাটের পেছনের প্রোফাইলার UI। এটি সরাসরি ট্রেস ডিরেক্টরি থেকে ডেটা পড়ে এবং একটি ওয়েব ইন্টারফেস প্রদান করে, যেখানে ট্রেস ভিউয়ার, মেমরি ভিউয়ার, গ্রাফ ভিউয়ার এবং অপারেশন পরিসংখ্যান দেখা যায়।

পডের ভিতরে এটি শুরু করুন:

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

সেলটি XProf যে URL-এ পরিষেবা দিচ্ছে, তার প্রসেস আইডি এবং এটিকে বন্ধ করার kill কমান্ডটি প্রিন্ট করে। ওই PID-টি লিখে রাখুন — ক্লিন আপ ধাপে আপনার এটি প্রয়োজন হবে।

আপনার ল্যাপটপ থেকে দর্শকের কাছে পৌঁছান।

ডিফল্টরূপে, পডের ভিতরে থাকা কোনো কিছুই আপনার মেশিন থেকে অ্যাক্সেসযোগ্য নয়। kubectl port-forward ক্লাউড শেল থেকে পডের ভিতরের একটি পোর্টে একটি টানেল খোলে এবং ক্লাউড শেলের ওয়েব প্রিভিউ সেই টানেলটি আপনার ব্রাউজারে প্রকাশ করে।

একটি নতুন ক্লাউড শেল ট্যাব খুলুন — কমান্ডটি আপনি বন্ধ না করা পর্যন্ত ফোরগ্রাউন্ডে চলতে থাকে এবং এই সময়ের মধ্যে আপনার নোটবুকটি পড-এ চলতে থাকে:

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

এরপর ক্লাউড শেল টুলবারে থাকা ‘Web Preview - Preview on port 8080’-এ ক্লিক করুন এবং আপনি UI অ্যাক্সেস করতে সক্ষম হবেন।

ট্রেসটি পড়ুন

একবার ওয়েব প্রিভিউ ট্যাবটি XProf লোড করলে:

  1. রানস ড্রপডাউন থেকে রানটি বেছে নিন।
  2. টুলস - ট্রেস_ভিউয়ার খুলুন।
  3. train_step স্প্যানগুলো খুঁজুন — এগুলো হলো পূর্ববর্তী ধাপের StepTraceAnnotation নাম।
  4. কম্পাইল স্প্যান, জিপিইউ গ্যাপ, কপি এবং কার্নেল অ্যাক্টিভিটির দিকে নজর রাখুন।

ট্রেস রিডিং চেকলিস্ট

নিচের সারণিতে দেখানো হয়েছে, চার্টের ভিজ্যুয়াল প্যাটার্ন কীভাবে পরবর্তী পদক্ষেপটি ট্রিগার করে:

চাক্ষুষ সংকেত

এর সম্ভবত মানে হলো

এরপর কী চেষ্টা করা যায়

প্রথম ধাপের আগে দীর্ঘ কম্পাইল স্প্যান।

সাধারণ প্রথম-কল JIT সংকলন

মাপার আগে গরম করে নিন।

ধাপগুলির মধ্যে স্প্যানগুলি সংকলন করুন

পুনঃসংকলন

কোডল্যাব ২ থেকে আকার, ডেটার ধরণ বা স্ট্যাটিক আর্গুমেন্ট পরিবর্তন পরীক্ষা করুন।

GPU সারিগুলিতে সাদা ফাঁক রয়েছে

হোস্ট জিপিইউ-তে ডেটা সরবরাহ করছে না।

ডেটা লোডিং, প্রিন্টিং, float(loss) এবং নামপাই কনভার্সন খুঁজুন।

অনেক ছোট ছোট দানা

ওভারহেড বা খুব ছোট সংকলিত অঞ্চল চালু করুন

JIT একটি বৃহত্তর ফাংশন; ব্যাচ ওয়ার্ক।

ধাপগুলির মধ্যে Memcpy কার্যকলাপ

হোস্ট-ডিভাইস স্থানান্তর

ডিভাইসে মেট্রিক্স রাখুন; কম ঘন ঘন লগ করুন; ট্রান্সফার গার্ড ব্যবহার করুন।

উচ্চ শিখর স্মৃতি

অ্যাক্টিভেশন বা অস্থায়ী বাফার প্রাধান্য পেতে পারে

memory_viewer খুলুন এবং ছোট ব্যাচ চেষ্টা করুন।

আরও তিনটি এন্ট্রি পয়েন্ট উল্লেখ করার মতো, যদিও এই কোডল্যাবে সেগুলো ব্যবহার করা হয়নি: start_trace() এবং stop_trace() যা ` with jax.profiler.trace(...) ব্লকের মধ্যে জায়গা হয় না এমন প্রোগ্রাম্যাটিক ট্রেস রিজিয়নের জন্য ব্যবহৃত হয়; দীর্ঘক্ষণ ধরে চলা জবগুলোর প্রোফাইলিং করার জন্য start_server()python -m jax.collect_profile ; এবং `Perfetto UI`-তে ট্রেস খোলার জন্য `Perfetto export`। 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 এর প্রথম লাইনটি প্রিন্ট করার মাধ্যমে সেলটির কাজ শেষ হওয়া উচিত।

ব্যাচের আকার এবং থ্রুপুট

ছোট ব্যাচ প্রায়শই জিপিইউ-কে পর্যাপ্ত সমান্তরাল কাজ দেয় না। সাধারণত ব্যাচের আকার বাড়ার সাথে সাথে থ্রুপুট উন্নত হয়, কিন্তু জিপিইউ-এর কাজের চাপ পূর্ণ হয়ে গেলে বা মেমোরিই সর্বোচ্চ সীমায় পৌঁছে গেলে তা স্থিতিশীল হয়ে যায়।

ব্যাচ সাইজও মেমরিকে প্রভাবিত করে। নিচের সুইপটিতে এই ফরোয়ার্ড পাসের ব্যাচ-শেপড বাফারগুলোর (ইনপুট, হিডেন অ্যাক্টিভেশন এবং আউটপুট) একটি সাধারণ আনুমানিক হিসাব অন্তর্ভুক্ত রয়েছে। প্রকৃত প্রশিক্ষণে এর চেয়ে বেশি ব্যবহৃত হয়, কারণ গ্রেডিয়েন্ট এবং অস্থায়ী বাফারগুলোও গণনা করা হয়।

@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,
)

আপনি সাতটি সারি এবং একটি বার চার্ট সহ একটি টেবিল পাবেন। ms/call নয়, examples/sec কলামটি দেখুন: প্রথম কয়েকটি ব্যাচ সাইজের ক্ষেত্রে থ্রুপুট দ্রুত বাড়ার কথা এবং তারপর জিপিইউ স্যাচুরেট হয়ে গেলে তা স্থিতিশীল হয়ে যাওয়ার কথা। ms/call কলামটি পুরোটা পথ জুড়েই বাড়তে থাকে, যা প্রত্যাশিত।

৯. জিপিইউ মেমরি প্রেসার পরীক্ষা করুন

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 মেমরি এর একটি ক্ষুদ্র ভগ্নাংশ হওয়া উচিত।

JAX ইম্পোর্ট করার আগে মেমরি সেটিংস অবশ্যই সেট করতে হবে। নোটবুকের ক্ষেত্রে, এর অর্থ হলো কার্নেল চালু হওয়ার আগে সেগুলো সেট করা এবং তারপর কার্নেলটি পুনরায় চালু করা।

পরিবর্তনশীল

উদাহরণ

কখন ব্যবহার করবেন

XLA_PYTHON_CLIENT_MEM_FRACTION

0.50

আপনি একটি জিপিইউ শেয়ার করেন এবং চান 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")

তিনটিই সম্ভবত প্রিন্ট হবে। যার অর্থ আপনি ডিফল্ট সেটিংসে আছেন: প্রি-অ্যালোকেশন চালু, ৭৫% ফ্র্যাকশন।

১০. Nsight Systems-এর সাহায্যে একটি CUDA টাইমলাইন ক্যাপচার করুন।

JAX-এর জন্য XProf হলো সঠিক প্রথম প্রোফাইলার। Nsight Systems হলো দ্বিতীয় ভিউ, যা CUDA টাইমলাইনের প্রয়োজনে ব্যবহার করা যায়: যেমন স্ট্রিম, CUDA API কল, কার্নেল, মেমকপি, cuBLAS ও cuDNN কল এবং অবশেষে NCCL।

কার্যপ্রবাহটির চারটি অংশ রয়েছে:

  1. কাজের চাপ একটি সংক্ষিপ্ত স্ক্রিপ্টে রাখুন।
  2. আপনার প্রয়োজনীয় ধাপগুলোর চারপাশে NVTX রেঞ্জ যোগ করুন।
  3. nsys profile দিয়ে স্ক্রিপ্টটি চালান।
  4. Nsight Systems GUI ব্যবহার করে .nsys-rep ফাইলটি খুলুন।

ক্যাপচার স্ক্রিপ্টটি লিখুন

সরাসরি একটি নোটবুকের প্রোফাইলিং করা সহজ নয়, তাই নিচের এই কোডটি তার পরিবর্তে একটি ছোট স্বতন্ত্র স্ক্রিপ্ট লেখে। এটি একবার ওয়ার্ম-আপ করে, তারপর ক্যাপচার করার যোগ্য অঞ্চলটি চিহ্নিত করতে 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}")

সেলটি যে দুটি পাথ ব্যবহার করবে তা প্রিন্ট করে। এখনো কিছু রান হয়নি।

রান এনসাইট সিস্টেমস

নিচের কোডটি শুধুমাত্র তখনই ডেটা সংগ্রহ শুরু করে যখন 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 আবার কম্পাইল করে। এটি শেষ হলে আপনি রিপোর্টের পাথ ও সাইজ এবং একটি ডাউনলোড লিঙ্ক পাবেন। এই লিঙ্কটি একটি ডেটা ইউআরএল, তাই এটি আপনার ব্রাউজারের নোটবুক থেকেই সরাসরি কাজ করে — কোনো পোর্ট-ফরোয়ার্ডের প্রয়োজন নেই।

এনসাইট টাইমলাইনটি পড়ুন

GUI-তে রিপোর্টটি খুলতে:

  1. developer.nvidia.com/nsight-systems থেকে আপনার লোকাল মেশিনে Nsight Systems ইনস্টল করুন। এটি বিনামূল্যে পাওয়া যায় এবং Linux, macOS ও Windows-এ চলে।
  2. উপরের লিঙ্কটি ব্যবহার করে তৈরি হওয়া .nsys-rep রিপোর্টটি ডাউনলোড করুন।
  3. nsys-ui চালু করুন এবং File - Open ব্যবহার করে ফাইলটি খুলুন, অথবা এটিকে উইন্ডোতে ড্র্যাগ অ্যান্ড ড্রপ করুন।

এই সারিগুলো দিয়ে শুরু করুন:

  • NVTX : profile_region এবং train_step_* খুঁজুন।
  • CUDA GPU সারি : কার্নেল কভারেজ এবং সাদা ফাঁক পরীক্ষা করুন, যেখানে একটি সাদা ফাঁক মানে GPU-এর নিষ্ক্রিয় সময়।
  • CUDA API : cudaLaunchKernel , cudaMemcpy* , এবং সিনক্রোনাইজেশন কলগুলো খুঁজুন।
  • cuBLAS / cuDNN : ট্রেস করলে এখানে matmul, convolution, এবং attention লাইব্রেরির কলগুলো দেখা যায়।
  • ওএস রানটাইম ( 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"])

আপনার তিনটি পর্যন্ত টেবিল পাওয়া উচিত, যার প্রতিটিতে সময় অনুযায়ী শীর্ষ আটটি সারি দেখানো থাকবে। প্রথমে কার্নেল সারাংশটি পড়ুন: এটি মোট সময় অনুযায়ী জিপিইউ কার্নেলগুলোকে র‍্যাঙ্ক করে, যা আপনাকে বলে দেয় ডিভাইসটি আসলে কোথায় তার সময় ব্যয় করেছে। সেটিকে মেমরি-অপারেশন টেবিলের সাথে তুলনা করুন — যদি কপিগুলো সময়ের জন্য কার্নেলের সাথে প্রতিযোগিতা করে, তাহলে আপনি আগের সেই হোস্ট-ট্রান্সফার সমস্যায় ফিরে যাবেন।

১১. পরিষ্কার করুন

আপনি জিপিইউ নোডে কিছু চালাচ্ছেন কি না, তা নির্বিশেষে এটি বিল করে, তাই এই ধাপটি এড়িয়ে যাবেন না।

প্রথমে পডের ভিতরে শুরু করা ভিউয়ারদের থামান: kill চালান। kill XProf এবং TensorBoard সেলগুলো যে কমান্ডগুলো প্রিন্ট করেছে, সেগুলো চালু থাকা অবস্থায় যেকোনো Cloud Shell ট্যাবে Ctrl+C চাপুন এবং kubectl port-forward চালান।

লোডব্যালেন্সার এবং পার্সিস্টেন্ট ভলিউম সহ জুপিটার ওয়ার্কলোডটি মুছে ফেলুন:

kubectl delete -f deploy/jupyter.yaml

ক্লাস্টার, নোড পুল, ভিপিসি এবং সার্ভিস অ্যাকাউন্ট ধ্বংস করুন:

cd terraform
terraform destroy

নির্দেশিত হলে yes টাইপ করুন, তারপর নিশ্চিত করুন যে পিছনে কিছু ফেলে রাখা হয়নি:

gcloud container clusters list
gcloud compute instances list

এই প্রজেক্টের জন্য উভয়ই খালি থাকা উচিত। যদি আপনি শুধু এই সিরিজের জন্য একটি প্রজেক্ট তৈরি করে থাকেন, তাহলে আপনি এর পরিবর্তে ক্লাউড কনসোল থেকে পুরো প্রজেক্টটি মুছে ফেলতে পারেন।

১২. অভিনন্দন

আপনি time.perf_counter পরিমাপ থেকে শুরু করে CUDA টাইমলাইন পর্যন্ত একটি JAX প্রশিক্ষণ ধাপের সম্পূর্ণ প্রোফাইলিং করেছেন।

আপনি যা শিখেছেন

  • block_until_ready() ব্যবহার করে কীভাবে নির্ভুলভাবে JAX-এর সময় পরিমাপ করা যায়, এবং কেন শুধুমাত্র ডিসপ্যাচ-এর পরিমাপ অর্থহীন।
  • কীভাবে ফার্স্ট-কল কম্পাইলেশন টাইমকে ক্যাশড এক্সিকিউশন টাইম থেকে আলাদা করা যায়, এবং কীভাবে একটি নতুন ইনপুট শেপ রিকম্পাইল করতে বাধ্য করে
  • jax.profiler.trace ব্যবহার করে কীভাবে একটি ট্রেস ক্যাপচার করবেন এবং StepTraceAnnotationTraceAnnotation দিয়ে সেটিকে লেবেল করবেন
  • ক্লাউড শেল পোর্ট-ফরোয়ার্ডের মাধ্যমে কীভাবে এক্সপ্রফ-এ এবং বিকল্প ফ্রন্ট এন্ড হিসেবে টেনসরবোর্ড-এ সেই ট্রেসটি খোলা যায়।
  • ধীরগতির সাধারণ ধরণগুলো কীভাবে চিনবেন: রিকম্পাইলেশন, অতিরিক্ত ছোট ছোট অপারেশন, হোস্ট-ডিভাইস ট্রান্সফার, অদক্ষ ব্যাচ সাইজ এবং মেমরি প্রেসার।
  • memory_stats() ব্যবহার করে কীভাবে GPU মেমরি পরীক্ষা করা যায়, এবং স্টার্টআপের আগে XLA_PYTHON_CLIENT_MEM_FRACTION কীভাবে JAX-এর রিজার্ভেশন পরিবর্তন করে।
  • NVTX রেঞ্জ ব্যবহার করে Nsight Systems-এর সাহায্যে কীভাবে একটি CUDA টাইমলাইন ক্যাপচার ও রিড করা যায়, এবং nsys stats ব্যবহার করে কীভাবে সেটির সারসংক্ষেপ করা যায়।

পরবর্তী পদক্ষেপ

  • কোডল্যাব ৪: JAX, Optax, এবং Fashion-MNIST ব্যবহার করে GPU-তে একটি মডেলকে প্রশিক্ষণ দিন, যেখানে আপনি আসল ডেটার উপর একটি বাস্তব প্রশিক্ষণ চক্রের অভিজ্ঞতা লাভ করবেন এবং এই ল্যাবের প্রোফাইলিং অভ্যাসগুলো আপনার মধ্যে আগে থেকেই বিদ্যমান থাকবে।
  • nsys_cmd তে --gpu-metrics-devices=cuda-visible যোগ করে Nsight ক্যাপচারটি পুনরায় চালান, এবং টাইমলাইনে দেখা কার্নেল কভারেজের সাথে GPU ইউটিলাইজেশন কাউন্টারগুলো তুলনা করুন।
  • কার্নেল পুনরায় চালু করার আগে XLA_PYTHON_CLIENT_MEM_FRACTION=0.50 সেট করুন, তারপর memory_stats() সেলটি পুনরায় চালান এবং limit কলামটি কীভাবে পরিবর্তিত হয়েছে তা পরীক্ষা করুন।

রেফারেন্স নথি