آموزش یک مدل روی GPU با JAX، Optax و Fashion-MNIST

۱. مقدمه

مسیر یادگیری جکس روی پردازنده گرافیکی. آزمایشگاه ۴: ساخت یک حلقه آموزشی ساده روی پردازنده گرافیکی.

در آزمایشگاه‌های کد قبلی، تأیید کردید که JAX می‌تواند GPU را ببیند، یاد گرفتید که چگونه jax.jit یک تابع را ردیابی و کامپایل می‌کند و از profiler برای دیدن اینکه GPU واقعاً چه کاری انجام می‌دهد استفاده کردید. اکنون قطعات در رایج‌ترین گردش کار JAX کنار هم قرار می‌گیرند: یک حلقه آموزشی.

در این مورد، شما یک MLP کوچک را روی Fashion-MNIST ، یک مجموعه داده طبقه‌بندی تصویر واقعی با ۶۰،۰۰۰ نمونه آموزشی و ۱۰،۰۰۰ نمونه آزمایشی، آموزش می‌دهید. هر نمونه یک تصویر خاکستری ۲۸x۲۸ از یک لباس است. در پایان، شما یک مرحله آموزشی Optax کامپایل شده، اعداد توان عملیاتی قابل اعتماد و یک ماتریس درهم‌ریختگی خواهید داشت که آن اعداد را به تصاویر واقعی متصل می‌کند.

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

  • Fashion-MNIST را دانلود کنید و آن را به دسته‌هایی با اندازه ثابت که روی GPU اجرا می‌شوند، تغییر شکل دهید.
  • یک MLP کوچک به عنوان PyTree از آرایه‌های JAX بسازید و یک تابع زیان اسکالر بنویسید.
  • گرادیان‌ها را با jax.grad و jax.value_and_grad محاسبه کنید، سپس مرحله را با jax.jit کامپایل کنید.
  • SGD دست‌نویس را با یک بهینه‌ساز Optax AdamW جایگزین کنید و یک حلقه آموزشی کوتاه اجرا کنید
  • اندازه‌گیری توان عملیاتی بر حسب مثال در ثانیه و مرتبط کردن آن با توکن در ثانیه
  • مدل را ارزیابی کنید، یک ماتریس درهم‌ریختگی رسم کنید و float32 با محاسبه bfloat16 مقایسه کنید.

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

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

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

مدل ذهنی مرحله آموزش

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

قطعه

چه کاری انجام می‌دهد؟

بررسی مبتدی

params

وزن‌های مدل به صورت PyTree ذخیره می‌شوند

همان ساختار درختی grads

batch

تصاویر و برچسب‌ها

برای جلوگیری از کامپایل مجدد، هر مرحله شکل یکسانی دارد

loss_fn

عبور رو به جلو به علاوه تلفات اسکالر

jax.grad به یک تابع زیان اسکالر نیاز دارد.

jax.value_and_grad

محاسبه‌ی همزمان تلفات و گرادیان‌ها

گرادیان‌ها با شکل پارامترها مطابقت دارند

optimizer.update

گرادیان‌ها را به به‌روزرسانی تبدیل می‌کند

حالت بهینه‌ساز فروشگاه‌های آدام

optax.apply_updates

پارامترهای بعدی را تولید می‌کند

پارامترها تغییرناپذیر هستند، بنابراین درخت جدید را برمی‌گرداند

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

۲. قبل از شروع

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

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

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

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

gcloud config set project <YOUR_PROJECT_ID>

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

دستور زیر را در Cloud Shell اجرا کنید. Codelab 1 هر دستور را با جزئیات توضیح می‌دهد، از جمله سهمیه GPU و الزامات منطقه.

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

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

cp terraform.tfvars.example terraform.tfvars

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

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

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

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

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

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

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

!pip install --quiet optax matplotlib

هر دو بسته معمولاً در کانتینر NVIDIA JAX عرضه می‌شوند، بنابراین این دستور pip معمولاً نیازی به عملیات ندارد.

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

JAX، Optax و چند تابع کمکی را وارد کنید. این سلول همچنین تأیید می‌کند که backend پیش‌فرض یک 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 به عنوان backend پیش‌فرض و لیستی از دستگاه‌های CUDA و به دنبال آن پردازنده گرافیکی (GPU) که ​​بقیه codelab از آن استفاده خواهد کرد را ببینید. تابع‌های کمکی show_table و show_bars جداول نتایج و نمودارهای میله‌ای را که در مراحل بعدی مشاهده خواهید کرد، رندر می‌کنند.

۳. بارگیری و بازرسی Fashion-MNIST

قبل از اینکه بتوانید چیزی را آموزش دهید، به داده‌هایی روی GPU نیاز دارید که در هر مرحله تغییر نکنند. در این مرحله، Fashion-MNIST دانلود می‌شود، بررسی می‌شود و به دسته‌هایی با اندازه ثابت تبدیل می‌شود که روی دستگاه قرار می‌گیرند.

مجموعه داده‌ها را دانلود کنید

Fashion-MNIST از همان فرمت فایل IDX مانند MNIST استفاده می‌کند. توابع کمکی زیر فایل‌های فشرده را دانلود می‌کنند، مجموع‌های کنترلی آنها را تأیید می‌کنند و تصاویر و برچسب‌ها را در آرایه‌های 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",
)

جدول باید ۶۰،۰۰۰ تصویر آموزشی و ۱۰،۰۰۰ تصویر آزمایشی را گزارش کند که هر کدام شکل (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()

شما باید یک جدول دو در پنج ببینید که در هر کلاس یک لباس قابل تشخیص دارد و هر عنوان با تصویر زیر آن مطابقت دارد.

آماده‌سازی دسته‌های GPU با اندازه ثابت

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

این سلول پیکسل‌ها را به [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 باید دستگاه CUDA انتخاب شده در سلول تنظیمات را نشان دهد.

۴. مدل را تعریف کنید و یک گام گرادیان بردارید

با دسته‌هایی روی دستگاه، به دو چیز نیاز دارید: مدلی که یک دسته را به لوجیت تبدیل کند، و یک اتلاف اسکالر که بتوانید آن را از هم متمایز کنید. این مرحله هر دو را می‌سازد، سپس یک گام گرادیان واحد را به صورت دستی انجام می‌دهد تا بتوانید دقیقاً ببینید jax.grad چه چیزی را برمی‌گرداند.

مدل و زیان را تعریف کنید

این مدل یک MLP کوچک است. یک مدل کانولوشن معمولاً برای تصاویر بهتر است، اما یک MLP مکانیک‌های مرحله آموزش را قابل مشاهده نگه می‌دارد: ضرب ماتریس، غیرخطی بودن، ضرب ماتریس، تلفات، گرادیان‌ها، به‌روزرسانی بهینه‌ساز.

مرحله‌ی رو به جلو، وزن‌ها و فعال‌سازی‌ها را به compute_dtype تبدیل می‌کند، سپس logits را قبل از از دست دادن به float32 تبدیل می‌کند. در حال حاضر، نوع داده‌ی محاسبه‌شده 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 و همگی روی پردازنده گرافیکی هستند. عنوان، تعداد کل مقادیر قابل آموزش را گزارش می‌دهد.

محاسبه یک گام گرادیان به صورت دستی

jax.grad و jax.value_and_grad به طور پیش‌فرض نسبت به اولین آرگومان تمایز قائل می‌شوند. در اینجا اولین آرگومان params است، بنابراین گرادیان همان ساختار درختی دیکشنری پارامتر را دارد.

وقتی فقط به گرادیان‌ها نیاز دارید jax.grad استفاده کنید. وقتی می‌خواهید مقدار loss و گرادیان‌ها را از یک مسیر رفت و برگشت داشته باشید، از 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 آن را به یک فایل اجرایی برای این شکل و نوع دسته‌ای کامپایل می‌کند.

اولین فراخوانی jitted شامل کامپایل می‌شود. زمان‌بندی زیر یک بار اجرا می‌شود و سپس اجرای cache شده را اندازه‌گیری می‌کند.

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

شما باید دو ردیف و یک نمودار دو میله‌ای ببینید، که در آن گام jitted به تعداد مثال/ثانیه بالاتری نسبت به توزیع پایتون می‌رسد. نسبت دقیق به پردازنده گرافیکی شما و شکل دسته بستگی دارد، اما نکته در ترتیب است: گام کامپایل شده همان کار را با سربار بسیار کمتر به ازای هر عملیات انجام می‌دهد.

۶. با یک بهینه‌ساز 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 است.

به الگوی ثبت وقایع توجه کنید. معیارها در طول هر مرحله روی دستگاه باقی می‌مانند و فقط هر چند مرحله به اعداد اعشاری پایتون تبدیل می‌شوند. دریافت زیان در هر تکرار به پایتون راحت است، اما میزبان را در هر تکرار با پردازنده گرافیکی (GPU) همگام‌سازی می‌کند.

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

شما باید یک جدول پنج در پنج ببینید که عناوین سبز رنگ در آن غالب هستند و تعداد انگشت‌شماری عنوان قرمز رنگ هم در آن دیده می‌شود. به عناوین قرمز رنگ نگاه کنید: بیشتر آنها باید ترکیبی از انواع لباس‌هایی باشند که در یک تصویر خاکستری ۲۸ در ۲۸ شبیه به هم به نظر می‌رسند.

ماتریس درهم‌ریختگی

تصاویری که ما استفاده می‌کنیم یک نمونه کوچک هستند. یک ماتریس درهم‌ریختگی نشان می‌دهد که مدل کدام کلاس‌ها را در کل مجموعه تست با هم ترکیب می‌کند. ردیف‌ها برچسب‌های واقعی و ستون‌ها برچسب‌های پیش‌بینی‌شده هستند و قطر به این معنی است که مدل معمولاً درست است.

@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 پشتیبانی‌شده استفاده کند. این به طور خودکار برای هر مدلی، به خصوص مدل‌های کوچک، سریع‌تر نیست، بنابراین به طور کلی: نوع داده را تغییر دهید، گرم کنید، اندازه‌گیری کنید.

در این نسخه ساده، پارامترها در float32 باقی می‌مانند. در طول انتقال رو به جلو، وزن‌ها و فعال‌سازی‌ها به نوع داده محاسباتی انتخاب شده تبدیل می‌شوند و logitها قبل از از دست دادن به 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 باقی می‌مانند.

بررسی کنید چه چیزی روی پردازنده گرافیکی (GPU) باقی مانده است

حلقه آموزش، پارامتر جدید و PyTrees با وضعیت بهینه‌ساز را برگرداند. آنها هنوز آرایه‌های JAX در GPU هستند. معیارها فقط زمانی به مقادیر پایتون تبدیل می‌شوند که شما آنها را به صراحت ثبت کنید.

در خطوط لوله ورودی واقعی، دسته‌ها اغلب از میزبان شروع می‌شوند. این اشکالی ندارد، اما کل یک دسته را در یک زمان منتقل کنید و از تبدیل‌هایی مانند 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 را نامگذاری کند. هیچ چیز در طول آموزش به طور مخفیانه به میزبان منتقل نشد، که این همان ویژگی است که شما قبل از افزایش مقیاس این حلقه در آزمایشگاه‌های کد بعدی می‌خواهید.

۹. تمیز کردن

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

kubectl delete -f deploy/jupyter.yaml

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

cd terraform
terraform destroy

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

gcloud container clusters list
gcloud compute instances list

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

۱۰. تبریک

شما از مفاهیم JAX منفرد به یک حلقه آموزش کامل GPU روی یک مجموعه داده واقعی منتقل شدید، و یک MLP را روی Fashion-MNIST از ابتدا تا انتها آموزش دادید.

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

  • نحوه ذخیره پارامترهای مدل به عنوان یک PyTree از آرایه‌های JAX و نگه داشتن دسته‌هایی با اندازه ثابت در GPU
  • چگونه یک تابع زیان اسکالر برای هدفی که می‌خواهید بهینه کنید، بنویسید؟
  • چه زمانی از jax.grad (فقط گرادیان‌ها) و چه زمانی از jax.value_and_grad (از دست دادن و گرادیان‌ها از یک گذر) استفاده کنیم؟
  • چگونه یک به‌روزرسانی Optax AdamW را درون یک گام آموزشی کامپایل‌شده با jax.jit قرار دهیم؟
  • چگونه می‌توان صادقانه میزان گذردهی را اندازه‌گیری کرد: ابتدا گرم کنید، قبل از متوقف کردن ساعت block_until_ready() استفاده کنید، و میزبان گیت، پشت یک وضعیت ثبت وقایع را می‌خواند.
  • چه ارتباطی بین examples/sec و tokens/sec وجود دارد، و چرا عدد tokens/sec در اینجا یک پیش‌بینی است نه یک اندازه‌گیری
  • چگونه پیش‌بینی‌ها و ماتریس درهم‌ریختگی را تجسم کنیم تا معیارها به نمونه‌های واقعی متصل شوند
  • چرا bfloat16 یک گزینه مفید برای اندازه‌گیری است و نه یک افزایش سرعت تضمین‌شده، و چرا پارامترها در float32 باقی می‌مانند؟

مراحل بعدی

  • Codelab 5: افزایش سرعت توجه در GPU با cuDNN و TransformerEngine که در آن همان نظم گام به گام کامپایل شده را در هسته‌های توجه که بر آموزش Transformer تسلط دارند، به کار خواهید گرفت.
  • BATCH_SIZE تغییر دهید و دوباره اجرا کنید. دسته‌های بزرگتر اغلب استفاده از GPU را بهبود می‌بخشند تا زمانی که حافظه به حد نهایی خود برسد و هر شکل جدید هزینه یک کامپایل مجدد را داشته باشد.
  • HIDDEN1 و HIDDEN2 را تغییر دهید. ضرب ماتریسی بیشتر معمولاً GPU را شلوغ‌تر می‌کند، بنابراین ببینید چه اتفاقی برای examples/sec می‌افتد.
  • log_every در train_many_steps تغییر دهید. ثبت وقایع بیشتر به معنای همگام‌سازی بیشتر میزبان است و عدد مربوط به توان عملیاتی باید این را نشان دهد.

اسناد مرجع