JAX, Optax, এবং Fashion-MNIST ব্যবহার করে GPU-তে একটি মডেলকে প্রশিক্ষণ দিন।

১. ভূমিকা

জিপিইউ-তে জ্যাক্স শেখার পথ। ল্যাব ৪: জিপিইউ-তে একটি সাধারণ ট্রেনিং লুপ তৈরি করা।

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

এই ধাপে আপনি Fashion-MNIST ডেটাসেটের উপর একটি ছোট MLP মেশিনকে প্রশিক্ষণ দেবেন। Fashion-MNIST হলো একটি বাস্তব ইমেজ-ক্লাসিফিকেশন ডেটাসেট, যেখানে ৬০,০০০ ট্রেনিং উদাহরণ এবং ১০,০০০ টেস্ট উদাহরণ রয়েছে। প্রতিটি উদাহরণ হলো একটি পোশাকের ২৮x২৮ আকারের গ্রেস্কেল ছবি। সবশেষে আপনার কাছে একটি সংকলিত Optax ট্রেনিং ধাপ, নির্ভরযোগ্য থ্রুপুট সংখ্যা এবং একটি কনফিউশন ম্যাট্রিক্স থাকবে, যা এই সংখ্যাগুলোকে বাস্তব ছবির সাথে সংযুক্ত করবে।

আপনি যা করবেন

  • Fashion-MNIST ডাউনলোড করুন এবং সেটিকে GPU-তে থাকা নির্দিষ্ট আকারের ব্যাচে রূপান্তর করুন।
  • JAX অ্যারের একটি PyTree হিসেবে একটি ছোট MLP তৈরি করুন এবং একটি স্কেলার লস ফাংশন লিখুন।
  • jax.grad এবং jax.value_and_grad ব্যবহার করে গ্রেডিয়েন্ট গণনা করুন, তারপর jax.jit দিয়ে ধাপটি কম্পাইল করুন।
  • হাতে লেখা SGD-কে Optax AdamW অপটিমাইজার দিয়ে প্রতিস্থাপন করুন এবং একটি সংক্ষিপ্ত ট্রেনিং লুপ চালান।
  • থ্রুপুটকে examples/sec এককে পরিমাপ করুন এবং এটিকে tokens/sec এর সাথে সম্পর্কিত করুন।
  • মডেলটি মূল্যায়ন করুন, একটি কনফিউশন ম্যাট্রিক্স অঙ্কন করুন, এবং float32bfloat16 এর মধ্যে তুলনা করুন।

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

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

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

প্রশিক্ষণ-ধাপ মানসিক মডেল

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

টুকরো

এটা যা করে

শিক্ষানবিস চেক

params

মডেলের ওজন PyTree হিসেবে সংরক্ষিত আছে

grads মতো একই বৃক্ষ কাঠামো

batch

ছবি এবং লেবেল

পুনঃসংকলন এড়াতে প্রতিটি ধাপে একই আকার।

loss_fn

ফরোয়ার্ড পাস প্লাস স্কেলার লস

jax.grad একটি স্কেলার লস প্রয়োজন।

jax.value_and_grad

ক্ষতি এবং গ্রেডিয়েন্ট একসাথে গণনা করে

গ্রেডিয়েন্টগুলি প্যারামিটারের আকারগুলির সাথে মেলে

optimizer.update

গ্রেডিয়েন্টকে আপডেটে রূপান্তর করে

অ্যাডাম অপ্টিমাইজারের অবস্থা সংরক্ষণ করে।

optax.apply_updates

পরবর্তী পরামিতিগুলি তৈরি করে

প্যারামিটারগুলো অপরিবর্তনীয়, তাই নতুন ট্রি-টি ফেরত দিন।

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

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

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

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

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

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

gcloud config set project <YOUR_PROJECT_ID>

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

ক্লাউড শেলে নিম্নলিখিতটি চালান। কোডল্যাব ১-এ জিপিইউ কোটা এবং জোনের প্রয়োজনীয়তা সহ প্রতিটি কমান্ড বিস্তারিতভাবে ব্যাখ্যা করা হয়েছে।

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 optax matplotlib

উভয় প্যাকেজই সাধারণত এনভিডিয়া জেএএক্স (NVIDIA JAX) কন্টেইনারে থাকে, তাই এই পিপ (pip) কমান্ডটি সাধারণত কোনো কাজ করে না।

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

JAX, Optax এবং কয়েকটি হেল্পার ইম্পোর্ট করুন। এই সেলটি আরও যাচাই করে যে ডিফল্ট ব্যাকএন্ডটি একটি GPU।

import gzip
import hashlib
import html
import math
import pathlib
import struct
import time
import urllib.request
from functools import partial

from IPython.display import HTML, display
import matplotlib.pyplot as plt
import numpy as np

import jax
import jax.numpy as jnp

try:
    import optax
except ModuleNotFoundError as exc:
    raise ModuleNotFoundError(
        "This lesson requires Optax. Install it with: pip install optax"
    ) from exc


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

print(f"JAX version:     {jax.__version__}")
print(f"Optax version:   {getattr(optax, '__version__', 'unknown')}")
print(f"Default backend: {jax.default_backend()}")
print(f"Devices:         {devices}")

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


CLASS_NAMES = np.array([
    "T-shirt/top",
    "Trouser",
    "Pullover",
    "Dress",
    "Coat",
    "Sandal",
    "Shirt",
    "Sneaker",
    "Bag",
    "Ankle boot",
])


def block_tree(tree):
    """Wait until a PyTree of JAX arrays is ready on device."""
    return jax.block_until_ready(tree)


def tree_l2_norm(tree):
    """L2 norm of all leaves in a PyTree treated as one long vector."""
    leaves = jax.tree_util.tree_leaves(tree)
    return jnp.sqrt(sum(jnp.sum(jnp.square(x)) for x in leaves))


def count_params(params):
    """Total number of scalar values across all leaves of a parameter PyTree."""
    return sum(x.size for x in jax.tree_util.tree_leaves(params))


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)
    parts = ["<div style='font-family: system-ui; max-width: 980px;'>"]
    if title:
        parts.append(f"<h4 style='margin: 0 0 8px 0;'>{html.escape(title)}</h4>")
    parts.append("<table style='border-collapse: collapse; width: 100%; font-size: 13px;'>")
    parts.append("<thead><tr>")
    for h, a in zip(headers, aligns):
        parts.append(
            f"<th style='text-align:{a}; border-bottom:1px solid #d0d7de; padding:6px;'>"
            f"{html.escape(str(h))}</th>"
        )
    parts.append("</tr></thead><tbody>")
    for row in rows:
        parts.append("<tr>")
        for cell, a in zip(row, aligns):
            parts.append(
                f"<td style='text-align:{a}; border-bottom:1px solid #eef1f4; padding:6px;'>"
                f"{html.escape(str(cell))}</td>"
            )
        parts.append("</tr>")
    parts.append("</tbody></table></div>")
    display(HTML("".join(parts)))


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

আপনি একটি JAX ভার্সন, একটি Optax ভার্সন, ডিফল্ট ব্যাকএন্ড হিসেবে gpu এবং CUDA ডিভাইসগুলোর একটি তালিকা দেখতে পাবেন, যার পরে কোডল্যাবের বাকি অংশে ব্যবহৃত GPU-টি দেখানো হবে। show_table এবং show_bars হেল্পারগুলো ফলাফলের টেবিল এবং বার চার্ট রেন্ডার করে, যা আপনি পরবর্তী ধাপগুলোতে দেখতে পাবেন।

৩. Fashion-MNIST লোড এবং পরীক্ষা করুন

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

ডেটাসেটটি ডাউনলোড করুন

Fashion-MNIST, MNIST-এর মতোই IDX ফাইল ফরম্যাট ব্যবহার করে। নিচের সহায়ক ফাংশনগুলো সংকুচিত ফাইলগুলো ডাউনলোড করে, সেগুলোর চেকসাম যাচাই করে এবং ছবি ও লেবেলগুলোকে NumPy অ্যারেতে পার্স করে।

এই রিপোজিটরিটিই হলো মূল উৎস। ডেটাসেট ফাইলগুলো স্থানীয়ভাবে ক্যাশ করা থাকে, তাই প্রথমবার চালানোর পর এই সেলটি দ্রুত কাজ করবে।

DATA_DIR = pathlib.Path.home() / ".cache" / "jax-course" / "fashion-mnist"
DATA_DIR.mkdir(parents=True, exist_ok=True)

FILES = {
    "train-images-idx3-ubyte.gz": "8d4fb7e6c68d591d4c3dfef9ec88bf0d",
    "train-labels-idx1-ubyte.gz": "25c81989df183df01b3e8a0aad5dffbe",
    "t10k-images-idx3-ubyte.gz": "bef4ecab320f06d8554ea6380940ec79",
    "t10k-labels-idx1-ubyte.gz": "bb300cfdad3c16e7a12a480ee83cd310",
}

PRIMARY_BASE_URL = "https://github.com/zalandoresearch/fashion-mnist/raw/master/data/fashion"

def md5sum(path):
    """Stream `path` in 1 MB chunks and return its MD5 hex digest."""
    digest = hashlib.md5()
    with open(path, "rb") as f:
        for chunk in iter(lambda: f.read(1024 * 1024), b""):
            digest.update(chunk)
    return digest.hexdigest()


def download_if_needed(filename, expected_md5):
    """Download `filename` if missing or its MD5 doesn't match."""
    path = DATA_DIR / filename
    if path.exists() and md5sum(path) == expected_md5:
        print(f"Using cached {filename}")
        return path

    urls = [f"{PRIMARY_BASE_URL}/{filename}"]
    last_error = None
    for url in urls:
        try:
            print(f"Downloading {filename} from {url}")
            urllib.request.urlretrieve(url, path)
            actual_md5 = md5sum(path)
            if actual_md5 != expected_md5:
                raise ValueError(f"MD5 mismatch: expected {expected_md5}, got {actual_md5}")
            return path
        except Exception as exc:
            last_error = exc
            if path.exists():
                path.unlink()
            print(f"  failed: {exc}")

    raise RuntimeError(f"Could not download {filename}") from last_error


def read_idx_images(path):
    """Parse the Fashion-MNIST IDX-3 image file at `path` and return a (N, rows, cols) uint8 array."""
    with gzip.open(path, "rb") as f:
        magic, num_images, rows, cols = struct.unpack(">IIII", f.read(16))
        assert magic == 2051, f"Unexpected image magic number {magic} in {path}"
        data = np.frombuffer(f.read(), dtype=np.uint8)
    return data.reshape(num_images, rows, cols)


def read_idx_labels(path):
    """Parse the IDX-1 label file at `path` and return a 1-D uint8 array of class indices."""
    with gzip.open(path, "rb") as f:
        magic, num_labels = struct.unpack(">II", f.read(8))
        assert magic == 2049, f"Unexpected label magic number {magic} in {path}"
        data = np.frombuffer(f.read(), dtype=np.uint8)
    return data.reshape(num_labels)


paths = {name: download_if_needed(name, checksum) for name, checksum in FILES.items()}

train_images = read_idx_images(paths["train-images-idx3-ubyte.gz"])
train_labels = read_idx_labels(paths["train-labels-idx1-ubyte.gz"])
test_images = read_idx_images(paths["t10k-images-idx3-ubyte.gz"])
test_labels = read_idx_labels(paths["t10k-labels-idx1-ubyte.gz"])

show_table(
    ["Split", "Images", "Image shape", "Labels"],
    [
        ("train", f"{len(train_images):,}", train_images.shape[1:], f"{len(train_labels):,}"),
        ("test", f"{len(test_images):,}", test_images.shape[1:], f"{len(test_labels):,}"),
    ],
    title="Fashion-MNIST loaded from IDX files",
)

টেবিলটিতে 60,000টি ট্রেনিং ইমেজ এবং 10,000টি টেস্ট ইমেজ রিপোর্ট করা উচিত, যার প্রতিটির আকার (28, 28) হবে এবং উভয় স্প্লিটের ইমেজের সমান সংখ্যক লেবেল থাকবে।

কয়েকটি উদাহরণ দেখুন

প্রশিক্ষণ শুরু করার আগে সবসময় কয়েকটি উদাহরণ দেখে নিন। এর মাধ্যমে অনেক বিরক্তিকর কিন্তু ব্যয়বহুল ত্রুটি ধরা পড়ে: যেমন ভুল লেবেল, ছবির ভুল দিকবিন্যাস, ভুল স্কেলিং, বা ভুলবশত ভুল ডেটাসেট লোড করা।

fig, axes = plt.subplots(2, 5, figsize=(10, 4))
for label, ax in enumerate(axes.flat):
    idx = np.flatnonzero(train_labels == label)[0]
    ax.imshow(train_images[idx], cmap="gray")
    ax.set_title(CLASS_NAMES[label], fontsize=10)
    ax.axis("off")
fig.suptitle("One Fashion-MNIST example per class")
fig.tight_layout()
plt.show()

আপনি একটি দুই বাই পাঁচ মাপের গ্রিড দেখতে পাবেন, যেখানে প্রতিটি শ্রেণীর জন্য একটি করে চেনা পোশাক থাকবে এবং প্রতিটি শিরোনামের সাথে তার নিচের ছবির মিল থাকবে।

নির্দিষ্ট আকারের জিপিইউ ব্যাচ প্রস্তুত করুন

এখন যেহেতু আপনি ডেটা যাচাই করে নিয়েছেন, আপনি সেগুলোকে ট্রেনিং লুপের কাঙ্ক্ষিত আকারে সাজিয়ে নিতে পারেন। JAX আরও বেশি কার্যকর হয় যখন প্রতিটি ট্রেনিং ধাপে একই আকার এবং ডেটার ধরন (dtypes) পাওয়া যায়।

সেলটি পিক্সেলগুলোকে [0, 1] পরিসরে স্বাভাবিক করে, প্রতিটি 28x28 ইমেজকে একটি 784-মানের ভেক্টরে পরিণত করে, ট্রেনিং সেটটিকে একবার শাফেল করে এবং ডেটাগুলোকে নির্দিষ্ট আকারের ব্যাচে রিসেপ করে। ব্যাচগুলোকে jax.device_put ব্যবহার করে একবার GPU-তে পাঠানো হয়, ফলে ট্রেনিং লুপটি শুধুমাত্র সেই অ্যারেগুলোকেই ইনডেক্স করে যেগুলো আগে থেকেই ডিভাইসে রয়েছে।

TRAIN_EXAMPLES = 60_000
TEST_EXAMPLES = 10_000
BATCH_SIZE = 512
INPUT_DIM = 28 * 28
NUM_CLASSES = 10


def prepare_images(images):
    """Cast uint8 images to float32 in [0, 1] and flatten each one into a 1-D feature vector."""
    images = images.astype(np.float32) / 255.0
    return images.reshape(images.shape[0], -1)


rng = np.random.default_rng(0)
train_perm = rng.permutation(len(train_images))[:TRAIN_EXAMPLES]

x_train = prepare_images(train_images[train_perm])
y_train = train_labels[train_perm].astype(np.int32)
x_test = prepare_images(test_images[:TEST_EXAMPLES])
y_test = test_labels[:TEST_EXAMPLES].astype(np.int32)


def make_fixed_batches(x, y, batch_size):
    """Trim trailing examples that don't fill a batch, reshape, and move to device."""
    usable = (len(x) // batch_size) * batch_size
    x = x[:usable].reshape(usable // batch_size, batch_size, x.shape[-1])
    y = y[:usable].reshape(usable // batch_size, batch_size)
    return jax.device_put(jnp.asarray(x), device), jax.device_put(jnp.asarray(y), device)


x_train_batches, y_train_batches = make_fixed_batches(x_train, y_train, BATCH_SIZE)
x_test_batches, y_test_batches = make_fixed_batches(x_test, y_test, BATCH_SIZE)
first_batch = (x_train_batches[0], y_train_batches[0])

show_table(
    ["Array", "Shape", "Dtype", "Devices"],
    [
        ("x_train_batches", x_train_batches.shape, x_train_batches.dtype, x_train_batches.devices()),
        ("y_train_batches", y_train_batches.shape, y_train_batches.dtype, y_train_batches.devices()),
        ("x_test_batches", x_test_batches.shape, x_test_batches.dtype, x_test_batches.devices()),
        ("y_test_batches", y_test_batches.shape, y_test_batches.dtype, y_test_batches.devices()),
    ],
    title="Fixed-size batches on GPU",
)

টেবিলে, প্রতিটি অ্যারে অবশ্যই float32 বা int32 হতে হবে, প্রতিটি শেপের শেষে ছবির জন্য 784 এবং লেবেলের জন্য 512 থাকতে হবে, এবং Devices কলামে আপনার setup সেলে নির্বাচিত CUDA ডিভাইসটি দেখানো হবে।

৪. মডেলটি সংজ্ঞায়িত করুন এবং একটি গ্রেডিয়েন্ট ধাপ নিন।

ডিভাইসে ব্যাচ ব্যবহার করার জন্য আপনার দুটি জিনিস প্রয়োজন: একটি মডেল যা একটি ব্যাচকে লজিটে রূপান্তর করে, এবং একটি স্কেলার লস যা আপনি ডিফারেনশিয়েট করতে পারেন। এই ধাপে উভয়ই তৈরি করা হয়, তারপর হাতে একটি গ্রেডিয়েন্ট ধাপ নেওয়া হয়, যাতে আপনি jax.grad ঠিক কী রিটার্ন করে তা দেখতে পারেন।

মডেল এবং ক্ষতি সংজ্ঞায়িত করুন

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

ফরওয়ার্ড পাস ওয়েট এবং অ্যাক্টিভেশনকে compute_dtype এ কাস্ট করে, তারপর লস গণনার আগে লজিটসকে আবার float32 এ কাস্ট করে। আপাতত `compute dtype` হলো float32 , কিন্তু পরে আপনি এটিকে bfloat16 এ পরিবর্তন করে দেখবেন কী পরিবর্তন হয়।

HIDDEN1 = 256
HIDDEN2 = 128
LEARNING_RATE = 3e-3


def init_mlp_params(key, input_dim=INPUT_DIM, hidden1=HIDDEN1, hidden2=HIDDEN2, num_classes=NUM_CLASSES):
    """Initialize a 3-layer MLP with He-style weight scaling and zero biases."""
    k1, k2, k3 = jax.random.split(key, 3)
    return {
        "w1": jax.random.normal(k1, (input_dim, hidden1), dtype=jnp.float32) * math.sqrt(2.0 / input_dim),
        "b1": jnp.zeros((hidden1,), dtype=jnp.float32),
        "w2": jax.random.normal(k2, (hidden1, hidden2), dtype=jnp.float32) * math.sqrt(2.0 / hidden1),
        "b2": jnp.zeros((hidden2,), dtype=jnp.float32),
        "w3": jax.random.normal(k3, (hidden2, num_classes), dtype=jnp.float32) * math.sqrt(2.0 / hidden2),
        "b3": jnp.zeros((num_classes,), dtype=jnp.float32),
    }


def mlp(params, x, compute_dtype=jnp.float32):
    """Forward pass: cast inputs/params to `compute_dtype`, two GELU hidden layers, then cast logits back to float32."""
    x = x.astype(compute_dtype)
    w1 = params["w1"].astype(compute_dtype)
    b1 = params["b1"].astype(compute_dtype)
    w2 = params["w2"].astype(compute_dtype)
    b2 = params["b2"].astype(compute_dtype)
    w3 = params["w3"].astype(compute_dtype)
    b3 = params["b3"].astype(compute_dtype)

    x = jax.nn.gelu(x @ w1 + b1)
    x = jax.nn.gelu(x @ w2 + b2)
    logits = x @ w3 + b3
    return logits.astype(jnp.float32)


def cross_entropy_loss(params, batch, compute_dtype=jnp.float32):
    """Scalar softmax cross-entropy loss."""
    x, y = batch
    logits = mlp(params, x, compute_dtype=compute_dtype)
    return optax.softmax_cross_entropy_with_integer_labels(logits, y).mean()


def loss_with_metrics(params, batch, compute_dtype=jnp.float32):
    """Same loss, but also returns batch accuracy in an aux dict."""
    x, y = batch
    logits = mlp(params, x, compute_dtype=compute_dtype)
    loss = optax.softmax_cross_entropy_with_integer_labels(logits, y).mean()
    accuracy = jnp.mean(jnp.argmax(logits, axis=-1) == y)
    return loss, {"accuracy": accuracy}


params = init_mlp_params(jax.random.key(1))
params = jax.device_put(params, device)

rows = []
for name, value in params.items():
    rows.append((name, value.shape, value.dtype, value.devices()))
show_table(["Parameter", "Shape", "Dtype", "Devices"], rows, title=f"MLP parameters: {count_params(params):,} trainable values")

সারণিটিতে ছয়টি লিফ তালিকাভুক্ত করা হয়েছে — তিনটি ওয়েট ম্যাট্রিক্স এবং তিনটি বায়াস ভেক্টর — সবগুলোই float32 এবং সবগুলোই GPU-তে অবস্থিত। শিরোনামে মোট প্রশিক্ষণযোগ্য মানের সংখ্যা উল্লেখ করা হয়েছে।

হাতে একটি গ্রেডিয়েন্ট ধাপ গণনা করুন

jax.grad এবং jax.value_and_grad ডিফল্টরূপে প্রথম আর্গুমেন্টের সাপেক্ষে পার্থক্য নির্ণয় করে। এখানে প্রথম আর্গুমেন্টটি হলো params , তাই গ্রেডিয়েন্টের ট্রি স্ট্রাকচারটি প্যারামিটার ডিকশনারির মতোই।

যখন শুধু গ্রেডিয়েন্ট প্রয়োজন, তখন jax.grad ব্যবহার করুন। কোনো ট্রেনিং স্টেপে, যখন আপনি একই ফরোয়ার্ড এবং ব্যাকওয়ার্ড পাস থেকে লস ও গ্রেডিয়েন্ট উভয়ই চান, তখন jax.value_and_grad ব্যবহার করুন।

নিচের কোডটি হলো সম্ভাব্য সবচেয়ে সরল প্রশিক্ষণ ধাপ: লস গণনা করা, গ্রেডিয়েন্ট গণনা করা, প্রতিটি প্যারামিটার থেকে একটি স্কেল করা গ্রেডিয়েন্ট বিয়োগ করা এবং একটি নতুন প্যারামিটার ট্রি ফেরত দেওয়া।

def sgd_step(params, batch):
    """One un-jitted SGD update and returns (new_params, loss, grads)."""
    loss, grads = jax.value_and_grad(cross_entropy_loss)(params, batch)
    new_params = jax.tree.map(lambda p, g: p - LEARNING_RATE * g, params, grads)
    return new_params, loss, grads


grads_only = jax.grad(cross_entropy_loss)(params, first_batch)
loss_value, grads = jax.value_and_grad(cross_entropy_loss)(params, first_batch)
loss_value, grads, grads_only = block_tree((loss_value, grads, grads_only))

rows = []
for name in params:
    rows.append((name, params[name].shape, grads[name].shape, grads[name].dtype))
show_table(["Leaf", "Param shape", "Grad shape", "Grad dtype"], rows, title="Gradient tree matches the parameter tree")

grad_difference = tree_l2_norm(jax.tree.map(lambda a, b: a - b, grads, grads_only))

print(f"loss before update: {float(loss_value):.4f}")
print(f"gradient L2 norm:   {float(tree_l2_norm(grads)):.4f}")
print(f"grad vs value_and_grad difference: {float(grad_difference):.6f}")

params_after_one, loss_after_one, _ = sgd_step(params, first_batch)
params_after_one, loss_after_one = block_tree((params_after_one, loss_after_one))
print(f"loss used for one SGD update: {float(loss_after_one):.4f}")

আউটপুটে তিনটি বিষয় যাচাই করতে হবে। প্রতিটি সারির Grad shape তার Param shape অনুরূপ হতে হবে, grad vs value_and_grad difference শূন্য বা এর খুব কাছাকাছি হতে হবে, কারণ উভয় রূপান্তর একই ডেরিভেটিভ গণনা করে এবং সর্বশেষ প্রিন্ট করা লসটি হলো আপডেটের আগে একই ব্যাচে গণনা করা লস।

sgd_step params পরিবর্তন না করে একটি নতুন প্যারামিটার ট্রি রিটার্ন করে, যা বাস্তবে 'পিওর ফাংশন' বলতে বোঝায়।

৫. jax.jit ব্যবহার করে প্রশিক্ষণ ধাপটি কম্পাইল করুন।

হাতে লেখা SGD ধাপটি সঠিক, কিন্তু একটি GPU ট্রেনিং লুপ চালানোর জন্য এটি সঠিক পদ্ধতি নয়। jit ছাড়া, পাইথন ক্রমাগত অনেকগুলো ছোট ছোট অপারেশন ডিসপ্যাচ করতে থাকে। jit ব্যবহার করলে, JAX পুরো ধাপটি একবার ট্রেস করে এবং XLA এই ব্যাচ শেপ ও dtype অনুযায়ী সেটিকে একটি এক্সিকিউটেবলে কম্পাইল করে।

প্রথম জিটেড কলে কম্পাইলেশন অন্তর্ভুক্ত থাকে। নীচের টাইমিংটি একবার ওয়ার্ম আপ করে, তারপর ক্যাশড এক্সিকিউশন পরিমাপ করে।

@jax.jit
def sgd_step_jit(params, batch):
    loss, grads = jax.value_and_grad(cross_entropy_loss)(params, batch)
    new_params = jax.tree.map(lambda p, g: p - LEARNING_RATE * g, params, grads)
    return new_params, loss


def batch_at(step):
    """Pick batch `step % num_batches`."""
    i = step % x_train_batches.shape[0]
    return x_train_batches[i], y_train_batches[i]


def time_loop(step_fn, params, steps):
    """Run `step_fn` for `steps` iterations and time it."""
    start = time.perf_counter()
    loss = None
    for step in range(steps):
        params, loss = step_fn(params, batch_at(step))
    params, loss = block_tree((params, loss))
    elapsed = time.perf_counter() - start
    return params, loss, elapsed


params_warm, loss_warm = sgd_step_jit(params, first_batch)
block_tree((params_warm, loss_warm))

EAGER_STEPS = 20
JIT_STEPS = 100

# sgd_step returns (params, loss, grads).
_, eager_loss, eager_elapsed = time_loop(lambda p, b: sgd_step(p, b)[:2], params, EAGER_STEPS)
_, jit_loss, jit_elapsed = time_loop(sgd_step_jit, params, JIT_STEPS)

eager_rate = EAGER_STEPS * BATCH_SIZE / eager_elapsed
jit_rate = JIT_STEPS * BATCH_SIZE / jit_elapsed

show_table(
    ["Mode", "Steps", "Final loss", "Elapsed seconds", "Examples/sec"],
    [
        ("Python dispatch", EAGER_STEPS, f"{float(eager_loss):.4f}", f"{eager_elapsed:.3f}", f"{eager_rate:,.0f}"),
        ("jitted step", JIT_STEPS, f"{float(jit_loss):.4f}", f"{jit_elapsed:.3f}", f"{jit_rate:,.0f}"),
    ],
    title="Cached training-step throughput",
    aligns=["left", "right", "right", "right", "right"],
)
show_bars([("Python dispatch", eager_rate), ("jitted step", jit_rate)], "Examples per second", "examples/s")

আপনি দুটি সারি এবং একটি দুই-বার চার্ট দেখতে পাবেন, যেখানে জিটেড স্টেপটি পাইথন ডিসপ্যাচের চেয়ে বেশি এক্সাম্পলস/সেকেন্ডে পৌঁছাবে। সঠিক অনুপাতটি আপনার জিপিইউ এবং ব্যাচ শেপের উপর নির্ভর করে, কিন্তু ক্রমটিই মূল বিষয়: কম্পাইলড স্টেপটি অনেক কম প্রতি-অপারেশন ওভারহেড সহ একই কাজ করে।

৬. একটি Optax অপ্টিমাইজার দিয়ে প্রশিক্ষণ নিন

গ্রেডিয়েন্ট দেখানোর জন্য র SGD যথেষ্ট, কিন্তু আসল ট্রেনিং লুপের জন্য মোমেন্টাম, ওয়েট ডিকে এবং শিডিউল প্রয়োজন হয়। এই ধাপে Optax যুক্ত করা হয়, একটি আসল ট্রেনিং লুপ চালানো হয় এবং ফলাফলটিকে একটি থ্রুপুট সংখ্যায় রূপান্তরিত করা হয়।

Optax অপ্টিমাইজার সেট আপ করুন

Optax , JAX-এর জন্য একটি গ্রেডিয়েন্ট প্রসেসিং এবং অপটিমাইজেশন লাইব্রেরি, আপনাকে মোমেন্টাম সহ SGD, Adam, AdamW, গ্রেডিয়েন্ট ক্লিপিং এবং লার্নিং-রেট শিডিউলের মতো কম্পোজেবল অপটিমাইজার প্রদান করে। একটি Optax অপটিমাইজারের দুটি গুরুত্বপূর্ণ মেথড রয়েছে: optimizer.init(params) অপটিমাইজার স্টেট তৈরি করে, যেমন Adam-এর মোমেন্টাম বাফার, এবং optimizer.update(grads, opt_state, params) গ্রেডিয়েন্টগুলোকে আপডেটে রূপান্তর করে এবং পরবর্তী অপটিমাইজার স্টেট রিটার্ন করে। এরপর optax.apply_updates(params, updates) পরবর্তী প্যারামিটার ট্রি রিটার্ন করে।

optimizer = optax.adamw(learning_rate=LEARNING_RATE, weight_decay=1e-4)
opt_state = optimizer.init(params)


@jax.jit
def train_step(params, opt_state, batch):
    (loss, metrics), grads = jax.value_and_grad(loss_with_metrics, has_aux=True)(params, batch)
    updates, opt_state = optimizer.update(grads, opt_state, params)
    params = optax.apply_updates(params, updates)
    metrics = {
        "loss": loss,
        "accuracy": metrics["accuracy"],
        "grad_norm": optax.global_norm(grads),
    }
    return params, opt_state, metrics


params_opt = init_mlp_params(jax.random.key(2))
params_opt = jax.device_put(params_opt, device)
opt_state = optimizer.init(params_opt)

params_opt, opt_state, metrics = train_step(params_opt, opt_state, first_batch)
params_opt, opt_state, metrics = block_tree((params_opt, opt_state, metrics))

show_table(
    ["Metric", "Value"],
    [
        ("loss", f"{float(metrics['loss']):.4f}"),
        ("accuracy", f"{100 * float(metrics['accuracy']):.1f}%"),
        ("gradient L2 norm", f"{float(metrics['grad_norm']):.4f}"),
    ],
    title="One compiled Optax training step",
    aligns=["left", "right"],
)

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

একটি ছোট প্রশিক্ষণ লুপ চালান।

একটি ধাপ কাজ করে, তাই আপনি এটি পুনরাবৃত্তি করতে পারেন। নীচের লুপটি একই আকারের ব্যাচগুলিতে একই কম্পাইল করা train_step কল করে। এটাই JAX প্রশিক্ষণের মূল প্যাটার্ন।

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

def train_many_steps(params, opt_state, steps=400, log_every=25):
    """Run a sized training loop, log a metric snapshot every `log_every` steps, and return history plus final state."""
    history = []
    start = time.perf_counter()
    metrics = None

    for step in range(steps):
        params, opt_state, metrics = train_step(params, opt_state, batch_at(step))

        if step % log_every == 0 or step == steps - 1:
            metrics = block_tree(metrics)
            history.append(
                {
                    "step": step,
                    "loss": float(metrics["loss"]),
                    "accuracy": float(metrics["accuracy"]),
                    "grad_norm": float(metrics["grad_norm"]),
                }
            )

    params, opt_state, metrics = block_tree((params, opt_state, metrics))
    elapsed = time.perf_counter() - start
    return params, opt_state, history, elapsed, metrics


params_train = init_mlp_params(jax.random.key(3))
params_train = jax.device_put(params_train, device)
# Fresh optimizer state, paired with this fresh `params_train`.
opt_state = optimizer.init(params_train)

params_train, opt_state, _ = train_step(params_train, opt_state, first_batch)
block_tree((params_train, opt_state))

TRAIN_STEPS = 400
params_train, opt_state, history, elapsed, final_metrics = train_many_steps(
    params_train, opt_state, steps=TRAIN_STEPS, log_every=25
)
examples_per_sec = TRAIN_STEPS * BATCH_SIZE / elapsed

show_table(
    ["Step", "Loss", "Accuracy", "Grad norm"],
    [(h["step"], f"{h['loss']:.4f}", f"{100*h['accuracy']:.1f}%", f"{h['grad_norm']:.3f}") for h in history],
    title=f"Training metrics, {examples_per_sec:,.0f} examples/sec",
    aligns=["right", "right", "right", "right"],
)

steps = [h["step"] for h in history]
losses = [h["loss"] for h in history]
accuracies = [h["accuracy"] for h in history]

fig, ax1 = plt.subplots(figsize=(8, 4))
ax1.plot(steps, losses, marker="o", color="#0969da", label="loss")
ax1.set_xlabel("step")
ax1.set_ylabel("loss", color="#0969da")
ax1.tick_params(axis="y", labelcolor="#0969da")
ax1.grid(True, alpha=0.25)

ax2 = ax1.twinx()
ax2.plot(steps, accuracies, marker="s", color="#1a7f37", label="accuracy")
ax2.set_ylabel("batch accuracy", color="#1a7f37")
ax2.tick_params(axis="y", labelcolor="#1a7f37")
ax2.set_ylim(0.0, 1.0)

fig.suptitle("Fashion-MNIST training curve")
fig.tight_layout()
plt.show()

আপনি একটি মেট্রিক্স টেবিল দেখতে পাবেন, যেখানে লগ করা প্রতিটি ধাপের জন্য একটি করে সারি থাকবে এবং একটি দ্বি-অক্ষীয় প্লট থাকবে, যেখানে ৪০০টি ধাপ জুড়ে ক্ষতির রেখাটি নিচে নামবে এবং ব্যাচ-সঠিকতার রেখাটি উপরে উঠবে।

৭. প্রশিক্ষিত মডেলটি মূল্যায়ন করুন।

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

স্থগিত পরীক্ষার তথ্যের উপর মূল্যায়ন করুন

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

মূল্যায়ন ফাংশনটি নির্দিষ্ট আকারের টেস্ট ব্যাচের উপর jax.vmap ব্যবহার করে। এটি কোডকে সংক্ষিপ্ত রাখে এবং JAX-কে সমস্ত ব্যাচের উপর একই প্রেডিকশন লজিক চালানোর সুযোগ দেয়।

@jax.jit
def evaluate_batches(params, x_batches, y_batches):
    """Vmap `loss_with_metrics` over every (x, y) batch and return the mean loss and accuracy."""
    def eval_one_batch(x, y):
        loss, metrics = loss_with_metrics(params, (x, y))
        return loss, metrics["accuracy"]

    losses, accuracies = jax.vmap(eval_one_batch)(x_batches, y_batches)
    return {"loss": jnp.mean(losses), "accuracy": jnp.mean(accuracies)}


test_metrics = evaluate_batches(params_train, x_test_batches, y_test_batches)
test_metrics = block_tree(test_metrics)

show_table(
    ["Split", "Loss", "Accuracy", "Examples evaluated"],
    [
        ("train batch", f"{float(final_metrics['loss']):.4f}", f"{100 * float(final_metrics['accuracy']):.1f}%", BATCH_SIZE),
        ("test", f"{float(test_metrics['loss']):.4f}", f"{100 * float(test_metrics['accuracy']):.1f}%", int(np.prod(y_test_batches.shape))),
    ],
    title="Evaluation after the short training run",
    aligns=["left", "right", "right", "right"],
)

টেস্ট রো-টি ফাইনাল ট্রেনিং ব্যাচের মতোই একই রেঞ্জে থাকা উচিত, উল্লেখযোগ্যভাবে খারাপ নয়। ' Examples evaluated ' কলামটি ছাঁটাই করা টেস্ট সেটটি দেখায়, সম্পূর্ণ ১০,০০০টি নয়, কারণ আপনার আগে তৈরি করা ব্যাচগুলো থেকে ২৭২টি অতিরিক্ত ডেটা বাদ পড়ে গেছে।

পূর্বাভাস কল্পনা করুন

এখন ভবিষ্যদ্বাণীগুলো দেখার পালা। সবুজ শিরোনামগুলো সঠিক ভবিষ্যদ্বাণী। লাল শিরোনামগুলো ভুল।

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

@jax.jit
def predict(params, x, compute_dtype=jnp.float32):
    """Return the predicted class index (argmax of the logits) for each row of `x`."""
    logits = mlp(params, x, compute_dtype=compute_dtype)
    return jnp.argmax(logits, axis=-1)


rng = np.random.default_rng(7)
sample_count = 25
sample_indices = rng.choice(len(test_images), size=sample_count, replace=False)
sample_pixels = test_images[sample_indices]
sample_x = prepare_images(sample_pixels)
sample_y = test_labels[sample_indices].astype(np.int32)

sample_x_device = jax.device_put(jnp.asarray(sample_x), device)
sample_pred = np.asarray(block_tree(predict(params_train, sample_x_device)))

fig, axes = plt.subplots(5, 5, figsize=(10, 10))
for ax, image, true_label, pred_label in zip(axes.flat, sample_pixels, sample_y, sample_pred):
    correct = int(true_label) == int(pred_label)
    ax.imshow(image, cmap="gray")
    ax.set_title(
        f"pred: {CLASS_NAMES[pred_label]}\ntrue: {CLASS_NAMES[true_label]}",
        fontsize=9,
        color="#1a7f37" if correct else "#d1242f",
    )
    ax.axis("off")
fig.suptitle("Sample Fashion-MNIST predictions")
fig.tight_layout()
plt.show()

আপনি একটি পাঁচ-বাই-পাঁচ গ্রিড দেখতে পাবেন, যেখানে বেশিরভাগই সবুজ শিরোনাম এবং হাতেগোনা কয়েকটি লাল শিরোনাম থাকবে। লাল শিরোনামগুলোর দিকে তাকান: এগুলোর বেশিরভাগই ২৮x২৮ মাপের একটি গ্রেস্কেল ছবিতে দেখতে একই রকম পোশাকের ধরন নিয়ে বিভ্রান্তি হওয়া উচিত।

বিভ্রান্তি ম্যাট্রিক্স

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

@jax.jit
def predict_batches(params, x_batches):
    """Run `predict` over every batch with vmap; output shape is (num_batches, batch_size)."""
    return jax.vmap(lambda x: predict(params, x))(x_batches)


test_pred = np.asarray(block_tree(predict_batches(params_train, x_test_batches))).reshape(-1)
test_true = np.asarray(y_test_batches).reshape(-1)

confusion = np.zeros((NUM_CLASSES, NUM_CLASSES), dtype=np.int32)
np.add.at(confusion, (test_true, test_pred), 1)
confusion_percent = confusion / confusion.sum(axis=1, keepdims=True)

fig, ax = plt.subplots(figsize=(8, 7))
im = ax.imshow(confusion_percent, cmap="Blues", vmin=0.0, vmax=1.0)
ax.set_xticks(np.arange(NUM_CLASSES), CLASS_NAMES, rotation=45, ha="right")
ax.set_yticks(np.arange(NUM_CLASSES), CLASS_NAMES)
ax.set_xlabel("Predicted label")
ax.set_ylabel("True label")
ax.set_title("Fashion-MNIST confusion matrix")
fig.colorbar(im, ax=ax, fraction=0.046, pad=0.04, label="fraction of true class")

for i in range(NUM_CLASSES):
    for j in range(NUM_CLASSES):
        value = confusion_percent[i, j]
        if value >= 0.08 or i == j:
            ax.text(j, i, f"{100 * value:.0f}%", ha="center", va="center", fontsize=8, color="white" if value > 0.45 else "black")

fig.tight_layout()
plt.show()

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

৮. float32 এবং bfloat16 এর গণনা তুলনা করুন

শুরুতে ট্রেনিং লুপের জন্য ডিফল্ট হলো float32bfloat16 কম বিট ব্যবহার করে, তাই এটি মেমরি ট্র্যাফিক কমাতে পারে এবং সমর্থিত NVIDIA GPU-তে দ্রুততর হার্ডওয়্যার পাথ ব্যবহার করতে পারে। এটি সব মডেলের জন্য, বিশেষ করে ছোট মডেলগুলোর জন্য, স্বয়ংক্রিয়ভাবে দ্রুততর নয়, তাই সাধারণভাবে: ডেটাটাইপ পরিবর্তন করুন, ওয়ার্ম আপ করুন এবং পরিমাপ করুন।

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

@partial(jax.jit, static_argnames=("compute_dtype",))
def train_step_mixed(params, opt_state, batch, compute_dtype=jnp.float32):
    """Same as `train_step`, but `compute_dtype` is a static argument."""
    (loss, metrics), grads = jax.value_and_grad(loss_with_metrics, has_aux=True)(
        params, batch, compute_dtype=compute_dtype
    )
    updates, opt_state = optimizer.update(grads, opt_state, params)
    params = optax.apply_updates(params, updates)
    metrics = {
        "loss": loss,
        "accuracy": metrics["accuracy"],
        "grad_norm": optax.global_norm(grads),
    }
    return params, opt_state, metrics


def time_mixed_precision(compute_dtype, steps=100):
    """Fresh init + one warmup compile for this dtype, then time `steps` steps and report examples/sec."""
    params_mp = init_mlp_params(jax.random.key(10))
    params_mp = jax.device_put(params_mp, device)
    opt_state_mp = optimizer.init(params_mp)

    params_mp, opt_state_mp, metrics = train_step_mixed(
        params_mp, opt_state_mp, first_batch, compute_dtype=compute_dtype
    )
    block_tree((params_mp, opt_state_mp, metrics))

    start = time.perf_counter()
    for step in range(steps):
        params_mp, opt_state_mp, metrics = train_step_mixed(
            params_mp, opt_state_mp, batch_at(step), compute_dtype=compute_dtype
        )
    params_mp, opt_state_mp, metrics = block_tree((params_mp, opt_state_mp, metrics))
    elapsed = time.perf_counter() - start
    return {
        "dtype": str(jnp.dtype(compute_dtype)),
        "loss": float(metrics["loss"]),
        "accuracy": float(metrics["accuracy"]),
        "elapsed": elapsed,
        "examples_per_sec": steps * BATCH_SIZE / elapsed,
    }


MIXED_PRECISION_STEPS = 100
mp_results = [
    time_mixed_precision(jnp.float32, steps=MIXED_PRECISION_STEPS),
    time_mixed_precision(jnp.bfloat16, steps=MIXED_PRECISION_STEPS),
]

show_table(
    ["Compute dtype", "Final loss", "Accuracy", "Elapsed seconds", "Examples/sec"],
    [
        (
            r["dtype"],
            f"{r['loss']:.4f}",
            f"{100 * r['accuracy']:.1f}%",
            f"{r['elapsed']:.3f}",
            f"{r['examples_per_sec']:,.0f}",
        )
        for r in mp_results
    ],
    title="Mixed-precision timing after warmup",
    aligns=["left", "right", "right", "right", "right"],
)
show_bars([(r["dtype"], r["examples_per_sec"]) for r in mp_results], "Mixed-precision examples/sec", "examples/s")

সারি দুটি তুলনা করুন। উভয়েরই লস এবং অ্যাকুরেসি প্রায় একই হওয়া উচিত, কারণ প্যারামিটার এবং লস উভয় ক্ষেত্রেই float32 এ থাকে।

জিপিইউ-তে কী রয়ে গেছে তা পরীক্ষা করুন

ট্রেনিং লুপটি নতুন প্যারামিটার এবং অপটিমাইজার-স্টেট পাইট্রি (PyTrees) রিটার্ন করেছে। এগুলো জিপিইউ-তে এখনও জ্যাক্স (JAX) অ্যারে হিসেবেই আছে। মেট্রিকগুলো কেবল তখনই পাইথন ভ্যালুতে পরিণত হয়, যখন আপনি সেগুলোকে স্পষ্টভাবে লগ করেন।

বাস্তব ইনপুট পাইপলাইনে, ব্যাচগুলো প্রায়শই হোস্টে শুরু হয়। তাতে কোনো সমস্যা নেই, কিন্তু পুরো ব্যাচটি একবারে স্থানান্তর করুন এবং লুপের সক্রিয় অংশে np.asarray(loss) বা float(loss) এর মতো রূপান্তর এড়িয়ে চলুন।

def devices_in_tree(tree):
    """Set of devices that any JAX-array leaf in `tree` currently lives on."""
    devices = set()
    for leaf in jax.tree_util.tree_leaves(tree):
        if hasattr(leaf, "devices"):
            devices.update(leaf.devices())
    return devices


show_table(
    ["Object", "Where its arrays live"],
    [
        ("trained params", devices_in_tree(params_train)),
        ("optimizer state", devices_in_tree(opt_state)),
        ("training batches", x_train_batches.devices()),
        ("test batches", x_test_batches.devices()),
    ],
    title="Device placement check",
)

print(f"final training loss = {float(final_metrics['loss']):.4f}")

প্লেসমেন্ট টেবিলের প্রতিটি সারিতে একটি CUDA ডিভাইসের নাম থাকা উচিত। ট্রেনিং চলাকালীন কোনো কিছুই স্বয়ংক্রিয়ভাবে হোস্টে স্থানান্তরিত হয় না, যা পরবর্তী কোডল্যাবগুলোতে এই লুপটিকে বড় পরিসরে নিয়ে যাওয়ার আগে আপনার কাঙ্ক্ষিত বৈশিষ্ট্য।

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

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

kubectl delete -f deploy/jupyter.yaml

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

cd terraform
terraform destroy

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

gcloud container clusters list
gcloud compute instances list

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

১০. অভিনন্দন

আপনি স্বতন্ত্র JAX ধারণা থেকে সরে এসে একটি বাস্তব ডেটাসেটের উপর সম্পূর্ণ GPU প্রশিক্ষণ চক্র শুরু করেছেন, এবং Fashion-MNIST ডেটাসেটের উপর একটি MLP মডেলকে শুরু থেকে শেষ পর্যন্ত প্রশিক্ষণ দিয়েছেন।

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

  • মডেল প্যারামিটারগুলোকে কীভাবে JAX অ্যারের PyTree হিসেবে সংরক্ষণ করা যায় এবং GPU-তে নির্দিষ্ট আকারের ব্যাচ রাখা যায়
  • আপনি যে উদ্দেশ্যটি অপ্টিমাইজ করতে চান তার জন্য কীভাবে একটি স্কেলার লস ফাংশন লিখতে হয়
  • কখন jax.grad (শুধুমাত্র গ্রেডিয়েন্টের জন্য) এবং কখন jax.value_and_grad (একই পাস থেকে লস এবং গ্রেডিয়েন্টের জন্য) ব্যবহার করতে হবে
  • কীভাবে একটি একক jax.jit কম্পাইল করা ট্রেনিং স্টেপের মধ্যে Optax AdamW আপডেট যুক্ত করবেন
  • সততার সাথে থ্রুপুট পরিমাপ করার উপায়: প্রথমে ওয়ার্ম আপ করুন, ক্লক বন্ধ করার আগে block_until_ready() ব্যবহার করুন, এবং একটি লগিং কন্ডিশনের আড়ালে হোস্ট রিডকে গেট করুন।
  • examples/sec কীভাবে tokens/sec-এর সাথে সম্পর্কিত, এবং কেন এখানে tokens/sec সংখ্যাটি একটি পরিমাপের পরিবর্তে একটি প্রক্ষেপণ।
  • কীভাবে পূর্বাভাস এবং একটি কনফিউশন ম্যাট্রিক্সকে দৃশ্যমান করা যায়, যাতে মেট্রিকগুলো বাস্তব উদাহরণের সাথে সংযুক্ত হয়।
  • কেন bfloat16 একটি নিশ্চিত গতিবৃদ্ধির পরিবর্তে পরিমাপের জন্য একটি দরকারী বিকল্প, এবং কেন প্যারামিটারগুলো float32 তেই থাকে

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

  • কোডল্যাব ৫: cuDNN এবং TransformerEngine ব্যবহার করে GPU-তে অ্যাটেনশনের গতি বৃদ্ধি করুন, যেখানে আপনি অ্যাটেনশন কার্নেলগুলিতে সেই একই কম্পাইলড-স্টেপ শৃঙ্খলা প্রয়োগ করবেন যা ট্রান্সফরমার প্রশিক্ষণে প্রাধান্য দেয়।
  • BATCH_SIZE পরিবর্তন করে পুনরায় চালান। বড় ব্যাচগুলি প্রায়শই GPU-এর ব্যবহার উন্নত করে যতক্ষণ না মেমরি সীমায় পৌঁছায়, এবং প্রতিটি নতুন আকারের জন্য একটি পুনঃসংকলন প্রয়োজন হয়।
  • HIDDEN1 এবং HIDDEN2 পরিবর্তন করুন। বেশি ম্যাট্রিক্স গুণন সাধারণত GPU-কে আরও ব্যস্ত করে তোলে, তাই examples/sec-এর কী হয় সেদিকে খেয়াল রাখুন।
  • train_many_steps এর log_every পরিবর্তন করুন। যত বেশি লগিং হবে, হোস্ট সিনক্রোনাইজেশন তত বাড়বে এবং থ্রুপুট সংখ্যায় তা প্রতিফলিত হওয়া উচিত।

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