১. ভূমিকা

পূর্ববর্তী কোডল্যাবগুলোতে, আপনারা যাচাই করেছেন যে 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 এর সাথে সম্পর্কিত করুন।
- মডেলটি মূল্যায়ন করুন, একটি কনফিউশন ম্যাট্রিক্স অঙ্কন করুন, এবং
float32ওbfloat16এর মধ্যে তুলনা করুন।
আপনার যা যা লাগবে
- বিলিং সক্ষম একটি গুগল ক্লাউড প্রজেক্ট, এবং ওয়ার্কশপ ক্রেডিট অথবা জিপিইউ ব্যবহারের জন্য একটি রিজার্ভেশন।
- আপনার নির্বাচিত অঞ্চলে কমপক্ষে ২টি এনভিডিয়া এল৪ জিপিইউ-এর জন্য কোটা ( জিপিইউ কোটা কীভাবে চেক করবেন )
- কোডল্যাব ১ থেকে ৩ সম্পন্ন করা, অথবা একটি সমতুল্য CUDA-সক্ষম JAX GPU পরিবেশ।
- প্রথম Fashion-MNIST ডাউনলোডের জন্য পড থেকে ইন্টারনেট সংযোগ।
সম্পূর্ণ করতে আনুমানিক সময়: ৬০ মিনিট ।
প্রশিক্ষণ-ধাপ মানসিক মডেল
একটি JAX ট্রেনিং স্টেপ হলো একটি পিওর ফাংশন: এটি অ্যারে ইনপুট হিসেবে নেয়, নতুন অ্যারে আউটপুট হিসেবে দেয় এবং পুরোনো প্যারামিটারগুলোকে সরাসরি পরিবর্তন করে না। নিচের প্রতিটি অংশে একটি বিগিনার চেক রয়েছে, যা কোনো সমস্যা হলে আপনি প্রয়োগ করতে পারেন।
টুকরো | এটা যা করে | শিক্ষানবিস চেক |
| মডেলের ওজন PyTree হিসেবে সংরক্ষিত আছে | |
| ছবি এবং লেবেল | পুনঃসংকলন এড়াতে প্রতিটি ধাপে একই আকার। |
| ফরোয়ার্ড পাস প্লাস স্কেলার লস | |
| ক্ষতি এবং গ্রেডিয়েন্ট একসাথে গণনা করে | গ্রেডিয়েন্টগুলি প্যারামিটারের আকারগুলির সাথে মেলে |
| গ্রেডিয়েন্টকে আপডেটে রূপান্তর করে | অ্যাডাম অপ্টিমাইজারের অবস্থা সংরক্ষণ করে। |
| পরবর্তী পরামিতিগুলি তৈরি করে | প্যারামিটারগুলো অপরিবর্তনীয়, তাই নতুন ট্রি-টি ফেরত দিন। |
এই অংশগুলো একটি নির্দিষ্ট চার-পর্যায়ের চক্রে চলে, যা এক ধাপ অগ্রসর হওয়া থেকে শুরু করে গ্রেডিয়েন্ট আপডেট এবং পুনরাবৃত্তি পর্যন্ত বিস্তৃত।
২. শুরু করার আগে
আপনার প্রকল্প নির্বাচন করুন
গুগল ক্লাউড কনসোলে , বিলিং সক্ষম করা আছে এমন একটি প্রজেক্ট নির্বাচন করুন বা তৈরি করুন।
ওপেন ক্লাউড শেল
একটি ক্লাউড শেল সেশন শুরু করতে অ্যাক্টিভেট ক্লাউড শেল (কনসোলের উপরের ডানদিকে থাকা টার্মিনাল আইকন)-এ ক্লিক করুন, তারপর এটিকে আপনার প্রজেক্টে নির্দেশ করুন:
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:// খুলুন http:// , টোকেনটি পেস্ট করুন এবং /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 এর গণনা তুলনা করুন
শুরুতে ট্রেনিং লুপের জন্য ডিফল্ট হলো float32 । bfloat16 কম বিট ব্যবহার করে, তাই এটি মেমরি ট্র্যাফিক কমাতে পারে এবং সমর্থিত 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পরিবর্তন করুন। যত বেশি লগিং হবে, হোস্ট সিনক্রোনাইজেশন তত বাড়বে এবং থ্রুপুট সংখ্যায় তা প্রতিফলিত হওয়া উচিত।