۱. مقدمه

در آزمایشگاههای کد قبلی، تأیید کردید که 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 یک تابع خالص است: آرایهها را دریافت میکند، آرایههای جدید را برمیگرداند و پارامترهای قدیمی موجود را تغییر نمیدهد. هر بخش زیر یک بررسی مبتدی دارد که میتوانید در صورت بروز مشکل اعمال کنید.
قطعه | چه کاری انجام میدهد؟ | بررسی مبتدی |
| وزنهای مدل به صورت PyTree ذخیره میشوند | همان ساختار درختی |
| تصاویر و برچسبها | برای جلوگیری از کامپایل مجدد، هر مرحله شکل یکسانی دارد |
| عبور رو به جلو به علاوه تلفات اسکالر | |
| محاسبهی همزمان تلفات و گرادیانها | گرادیانها با شکل پارامترها مطابقت دارند |
| گرادیانها را به بهروزرسانی تبدیل میکند | حالت بهینهساز فروشگاههای آدام |
| پارامترهای بعدی را تولید میکند | پارامترها تغییرناپذیر هستند، بنابراین درخت جدید را برمیگرداند |
این قطعات در یک چرخه چهار مرحلهای ثابت اجرا میشوند، از مرحله رو به جلو گرفته تا بهروزرسانی گرادیانها و تکرار آنها.
۲. قبل از شروع
پروژه خود را انتخاب کنید
در کنسول گوگل کلود ، یک پروژه با قابلیت پرداخت فعال انتخاب یا ایجاد کنید.
پوسته ابری را باز کنید
برای شروع یک جلسه 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:// را باز کنید http:// وارد کنید، توکن را جایگذاری کنید و یک دفترچه یادداشت پایتون ۳ جدید در /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تغییر دهید. ثبت وقایع بیشتر به معنای همگامسازی بیشتر میزبان است و عدد مربوط به توان عملیاتی باید این را نشان دهد.