۱. مقدمه

در این آزمایشگاه کد، شما یک مبدل رمزگشای کوچک در JAX با Flax NNX تعریف میکنید، آن را روی متن شکسپیر در هر دو GPU روی گره خود آموزش میدهید، آن را با Orbax ذخیره و بازیابی میکنید و متن جدیدی از وزنهای آموزش دیده تولید میکنید.
این مدل عمداً کوچک ساخته شده است و دارای ۴ لایه، جاسازیهای ۲۵۶ بعدی و واژگان سطح بایت است، به طوری که در کمتر از یک دقیقه روی دو پردازنده گرافیکی L4 آموزش میبیند. معماری و الگوهای آموزشی همانهایی هستند که در مدلهای بسیار بزرگتر استفاده میشوند.
کاری که انجام خواهید داد
- با استفاده از
nnx.Embed،nnx.MultiHeadAttention،nnx.Linearوnnx.LayerNorm، یک ترانسفورماتور رمزگشا با Flax NNX تعریف کنید. - به عنوان backend مربوط به attention، کد علت و معلولی
jax.nn.dot_product_attentionرا اضافه کنید. - آموزش در سطح بایت TinyShakespeare با
nnx.Optimizerو Optax AdamW، ابتدا روی یک پردازنده گرافیکی و سپس روی همه آنها - توان عملیاتی را بر حسب توکن بر ثانیه اندازهگیری کنید و دو اجرا را با هم مقایسه کنید.
- ذخیره و بازیابی پارامترهای مدل با Orbax
StandardCheckpointer - تولید متن شکسپیر مانند از مدل آموزش دیده
آنچه نیاز دارید
- یک پروژه Google Cloud با قابلیت پرداخت صورتحساب، و اعتبار کارگاه یا رزرو شامل استفاده از GPU
- سهمیه حداقل ۲ پردازنده گرافیکی NVIDIA L4 در منطقه انتخابی شما ( نحوه بررسی سهمیه پردازنده گرافیکی )
- Codelabs 1 تا 6 تکمیل شده، یا یک محیط GPU معادل JAX با حداقل دو GPU
- دسترسی به اینترنت خروجی از طریق پاد، بنابراین اولین اجرا میتواند فایل متنی TinyShakespeare را دانلود کند.
زمان تخمینی برای تکمیل: ۷۰ دقیقه .
معماری که شما میسازید
این مدل مجموعهای از بلوکهای ترانسفورماتور است. هر بلوک دارای دو زیرلایه است که هر کدام در یک اتصال باقیمانده پیچیده شدهاند:
- توجه به خود - هر موقعیت به تمام موقعیتهای قبلی توجه میکند (ماسک علّی).
- شبکه پیشخور (FFN) - دو لایه خطی با فعالسازی GELU، که نمایش را گسترش داده و سپس فشرده میکند.
هر دو زیرلایه از pre-norm استفاده میکنند و LayerNorm قبل از زیرلایه اعمال میشود. pre-norm برای آموزش پایدارتر است و در ترانسفورماتورهای مدرن استاندارد است.
۲. قبل از شروع
پروژه خود را انتخاب کنید
در کنسول گوگل کلود ، یک پروژه با قابلیت پرداخت فعال انتخاب یا ایجاد کنید.
پوسته ابری را باز کنید
برای شروع یک جلسه 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 flax optax orbax-checkpoint matplotlib
GPU را تنظیم و تأیید کنید
JAX، Flax NNX، Optax و Orbax را وارد کنید و بررسی کنید که کانتینر چند پردازنده گرافیکی را میتواند ببیند.
import os
os.environ["LD_LIBRARY_PATH"] = "/usr/local/nvidia/lib64:" + os.environ.get(
"LD_LIBRARY_PATH", ""
)
import hashlib
import html
import math
import pathlib
import time
import urllib.request
import warnings
from IPython.display import HTML, display
import matplotlib.pyplot as plt
import numpy as np
warnings.filterwarnings("ignore", category=DeprecationWarning)
warnings.filterwarnings("ignore", message=".*ml_dtypes.*")
warnings.filterwarnings("ignore", message=".*JAX_PLATFORMS.*")
import jax
import jax.numpy as jnp
import optax
from flax import nnx
import orbax.checkpoint as ocp
from jax.sharding import Mesh, PartitionSpec as P, NamedSharding
devices = jax.devices()
gpu_devices = [d for d in devices if d.platform == "gpu"]
NUM_DEVICES = len(gpu_devices)
print(f"JAX version: {jax.__version__}")
print(f"Default backend: {jax.default_backend()}")
print(f"GPU devices: {gpu_devices}")
print(f"GPU count: {NUM_DEVICES}")
assert len(gpu_devices) >= 2, (
f"This lesson needs at least 2 GPUs. Found {len(gpu_devices)}. "
f"Available devices: {devices}"
)
def block_tree(tree):
return jax.block_until_ready(tree)
def show_table(headers, rows, title=None, aligns=None):
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):
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):,.0f} {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، gpu به عنوان backend پیشفرض، لیستی از دو دستگاه CUDA و GPU count: 2 را ببینید. توابع کمکی block_tree ، show_table و show_bars جداول و نمودارهای میلهای مورد استفاده در مراحل بعدی را رندر میکنند.
۳. دادههای TinyShakespeare را در سطح بایت آماده کنید
TinyShakespeare یک فایل متنی واحد با حجم حدود ۱ مگابایت است که شامل آثار شکسپیر است که به هم پیوستهاند. این آزمایشگاه کد از توکنسازی در سطح بایت استفاده میکند که در آن هر بایت از متن UTF-8 به یک توکن تبدیل میشود. این کار واژگان را در ۲۵۶ مقدار ممکن تثبیت میکند و هرگونه وابستگی به توکنساز را حذف میکند.
متن به دنبالههای غیرهمپوشانی به طول SEQ_LEN تقسیم میشود. هر دنباله یک نمونه آموزشی است و مدل یاد میگیرد که بایت بعدی را در هر موقعیت پیشبینی کند.
کد را اجرا کنید تا فایل دانلود شود، مجموع مقابلهای آن تأیید شود، به توالیهای آموزش و اعتبارسنجی تقسیم شود و مجموعه آموزش تغییر کند:
SHAKESPEARE_URL = "https://raw.githubusercontent.com/karpathy/char-rnn/master/data/tinyshakespeare/input.txt"
SHAKESPEARE_MD5 = "d015dc5942f9b2908e24d4827a3e7a5e"
DATA_DIR = pathlib.Path.home() / ".cache" / "jax-course"
DATA_DIR.mkdir(parents=True, exist_ok=True)
DATA_FILE = DATA_DIR / "tinyshakespeare.txt"
def md5sum(path):
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()
if DATA_FILE.exists() and md5sum(DATA_FILE) == SHAKESPEARE_MD5:
print("Using cached tinyshakespeare.txt")
else:
print(f"Downloading tinyshakespeare.txt")
urllib.request.urlretrieve(SHAKESPEARE_URL, DATA_FILE)
raw_text = DATA_FILE.read_text()
data = np.frombuffer(raw_text.encode("utf-8"), dtype=np.uint8).astype(np.int32)
print()
VOCAB_SIZE = 256
SEQ_LEN = 256
PER_DEVICE_BATCH = 32
num_sequences = len(data) // SEQ_LEN
data = data[: num_sequences * SEQ_LEN].reshape(num_sequences, SEQ_LEN)
num_train = int(0.9 * num_sequences)
train_data = data[:num_train]
val_data = data[num_train:]
rng = np.random.default_rng(0)
train_data = train_data[rng.permutation(num_train)]
def make_batches(data, batch_size):
usable = (len(data) // batch_size) * batch_size
return data[:usable].reshape(-1, batch_size, SEQ_LEN)
show_table(
["", "Value"],
[
("Total bytes", f"{len(raw_text):,}"),
("Vocabulary", f"{VOCAB_SIZE} (raw bytes)"),
("Sequence length", SEQ_LEN),
("Training sequences", f"{num_train:,}"),
("Validation sequences", f"{len(val_data):,}"),
],
title="TinyShakespeare — byte-level tokenization",
)
print()
print("Sample text (first 200 bytes):")
print(raw_text[:200])
شما باید یا عبارت Using cached tinyshakespeare.txt یا یک پیام دانلود را ببینید، و پس از آن جدولی از شکل مجموعه دادهها با کل بایتها، یک واژگان ۲۵۶ بایتی خام، طول توالی، و تعداد توالیهای آموزش و اعتبارسنجی و سپس ۲۰۰ بایت اول متن آمده است تا بتوانید ببینید مدل از چه چیزی یاد میگیرد.
سه ثابتی که در اینجا تنظیم شدهاند برای بقیهی بخش کد اهمیت دارند. VOCAB_SIZE برابر با ۲۵۶ است زیرا یک بایت ۲۵۶ مقدار ممکن دارد و SEQ_LEN برابر با ۲۵۶ موقعیت در هر مثال آموزشی است.
PER_DEVICE_BATCH برابر با ۳۲ است و برای هر دو حالت تک پردازنده گرافیکی و چند پردازنده گرافیکی ثابت میماند. همین ثابت بودن تعداد پردازندههای گرافیکی به ازای هر پردازنده گرافیکی است که باعث میشود مقایسه توان عملیاتی در آینده به مقایسهای با مقیاس ضعیف تبدیل شود.
۴. ترانسفورماتور را با Flax NNX تعریف کنید
Flax NNX یک API سادهشده برای شبکههای عصبی در JAX است. شما لایهها را به عنوان اشیاء پایتون تعریف میکنید که مقداردهی اولیه وزن و مسیر رو به جلوی خود را دارند. این مرحله هر قطعه NNX مورد نیاز مبدل را معرفی میکند، بنابراین هیچ تجربه قبلی NNX فرض نمیشود.
این مدل از چهار بلوک سازنده استفاده میکند:
-
nnx.Embedیک جدول جستجو است که یک اندیس توکن را به یک بردار نگاشت میکند. -
nnx.Linearیک ضرب ماتریسی متراکم به علاوه یک بایاس اختیاری است. -
nnx.LayerNormویژگیهای قبل از زیرلایههای attention و FFN را نرمالسازی میکند. -
nnx.MultiHeadAttentionپیشبینیهای Q/K/V، توجه و پیشبینی خروجی را مدیریت میکند.
توجه سببی را به اشتراک بگذارید
nnx.MultiHeadAttention حالتهای پنهان به شکل (B, T, D_MODEL) را دریافت میکند، Q، K و V را به صورت داخلی ایجاد میکند و آنها را به سرها تقسیم میکند. قلاب attention_fn فقط عملیات اصلی توجه را که پس از آن پیشبینیها اجرا میشود، کنترل میکند.
NNX آرگومانهای اختیاری به سبک Flax مانند dropout rng، dtype و precision را به attention_fn ارسال میکند. jax.nn.dot_product_attention این آرگومانها را نمیپذیرد، بنابراین wrapper زیر آنها را با یک catch-all میگیرد و فقط آنچه را که تابع JAX نیاز دارد، ارسال میکند.
D_MODEL = 256
NUM_HEADS = 4
FFN_DIM = 1024
NUM_LAYERS = 4
MAX_SEQ_LEN = 256
LR = 3e-4
WEIGHT_DECAY = 1e-4
def causal_sdpa(query, key, value, **_):
return jax.nn.dot_product_attention(query, key, value, is_causal=True)
تعریف بلوک و مدل
دو کلاس وجود دارد. TransformerBlock یک زیرلایه attention به علاوه یک زیرلایه FFN است و TinyTransformer num_layers آنها را بین جاسازیها و سر LM روی هم قرار میدهد. هر کلاس nnx.Module را ارثبری میکند و تمام لایههای آن را در __init__ ایجاد میکند.
class TransformerBlock(nnx.Module):
def __init__(self, d_model: int, num_heads: int, ffn_dim: int, rngs: nnx.Rngs):
self.ln1 = nnx.LayerNorm(d_model, rngs=rngs)
self.attn = nnx.MultiHeadAttention(
num_heads=num_heads,
in_features=d_model,
decode=False,
attention_fn=causal_sdpa,
rngs=rngs,
)
self.ln2 = nnx.LayerNorm(d_model, rngs=rngs)
self.fc_up = nnx.Linear(d_model, ffn_dim, rngs=rngs)
self.fc_down = nnx.Linear(ffn_dim, d_model, rngs=rngs)
def __call__(self, x):
x = x + self.attn(self.ln1(x))
h = jax.nn.gelu(self.fc_up(self.ln2(x)))
x = x + self.fc_down(h)
return x
ترتیب پیش-هنجار در __call__ قابل مشاهده است. x = x + self.attn(self.ln1(x)) قبل از attention نرمالسازی میکند و نتیجه را به جریان باقیمانده اضافه میکند، و شاخه FFN همین کار را با self.ln2 انجام میدهد.
class TinyTransformer(nnx.Module):
def __init__(
self,
vocab_size: int,
d_model: int,
num_heads: int,
ffn_dim: int,
num_layers: int,
max_seq_len: int,
rngs: nnx.Rngs,
):
self.token_embed = nnx.Embed(vocab_size, d_model, rngs=rngs)
self.pos_embed = nnx.Embed(max_seq_len, d_model, rngs=rngs)
self.blocks = nnx.List(
[
TransformerBlock(d_model, num_heads, ffn_dim, rngs=rngs)
for _ in range(num_layers)
]
)
self.final_norm = nnx.LayerNorm(d_model, rngs=rngs)
self.lm_head = nnx.Linear(d_model, vocab_size, use_bias=False, rngs=rngs)
def __call__(self, tokens):
B, T = tokens.shape
x = self.token_embed(tokens) + self.pos_embed(jnp.arange(T))
for block in self.blocks:
x = block(x)
x = self.final_norm(x)
return self.lm_head(x)
TinyTransformer.__call__ جاسازی توکن و جاسازی موقعیت را اضافه میکند، بلوکها را به ترتیب اجرا میکند، یک LayerNorm نهایی اعمال میکند و به logits های vocab_size پروژه میدهد.
نمونهسازی و بررسی مدل
ایجاد مدل با یک فراخوانی انجام میشود. nnx.Rngs تمام حالتهای تصادفی مورد نیاز برای مقداردهی اولیه پارامترها را مدیریت میکند. پس از ایجاد، میتوانید پارامترهای آن را بشمارید و یک گذر رو به جلو از طریق آن اجرا کنید.
model = TinyTransformer(
VOCAB_SIZE,
D_MODEL,
NUM_HEADS,
FFN_DIM,
NUM_LAYERS,
MAX_SEQ_LEN,
rngs=nnx.Rngs(0),
)
param_count = sum(x.size for x in jax.tree.leaves(nnx.state(model, nnx.Param)))
show_table(
["", "Value"],
[
("Architecture", f"Decoder-only transformer"),
("Layers", NUM_LAYERS),
("Model dimension", D_MODEL),
("Attention heads", f"{NUM_HEADS} (head dim = {D_MODEL // NUM_HEADS})"),
("FFN dimension", FFN_DIM),
("Vocabulary", f"{VOCAB_SIZE} (byte-level)"),
("Max sequence length", MAX_SEQ_LEN),
("Parameters", f"{param_count:,}"),
],
title="TinyTransformer",
)
logits = model(jnp.zeros((1, 16), dtype=jnp.int32))
print(f"Test forward pass: input (1, 16) \u2192 logits {logits.shape}")
شما باید جدولی را ببینید که معماری را توصیف میکند: ۴ لایه، بعد مدل ۲۵۶، ۴ سر توجه با بعد سر ۶۴، یک واژگان سطح بایت و یک تعداد پارامتر. خط آخر، مسیر رو به جلوی تست را گزارش میدهد، با ورودی شکل (1, 16) که لوجیتهایی با شکل (1, 16, 256) تولید میکند - یک توزیع روی مقادیر ۲۵۶ بایتی برای هر یک از ۱۶ موقعیت ورودی.
۵. مرحله آموزش NNX را بنویسید
در اینجا شما @nnx.jit برای مدیریت خودکار وضعیت ماژول NNX استفاده میکنید. این دستور ماژولها را به ساختار و آرایههایی برای کامپایل JIT تقسیم میکند، سپس آرایههای بهروزرسانیشده را دوباره ادغام میکند، بنابراین شما مرحله را طوری مینویسید که انگار ماژولها اشیاء معمولی پایتون هستند.
ضرر، پیشبینی توکن بعدی است. در هر موقعیت، مدل توکن بعدی را پیشبینی میکند، بنابراین شما logits[:, :-1] (پیشبینیها در موقعیتهای 0 تا T-2) را با tokens[:, 1:] (توکنهای واقعی در موقعیتهای 1 تا T-1) مقایسه میکنید.
@nnx.jit
def train_step(model, optimizer, tokens):
def loss_fn(model):
logits = model(tokens)
pred = logits[:, :-1].reshape(-1, VOCAB_SIZE)
target = tokens[:, 1:].reshape(-1)
return optax.softmax_cross_entropy_with_integer_labels(pred, target).mean()
loss, grads = nnx.value_and_grad(loss_fn)(model)
optimizer.update(model, grads)
return {"loss": loss, "perplexity": jnp.exp(loss)}
حلقهی اطراف آن مرحله یک بار فعال میشود تا کامپایل زمانبندیشده نباشد، سپس تعداد ثابتی از مراحل را اجرا میکند و زمان سپریشده را به توکن/ثانیه تبدیل میکند.
def train_loop(model, optimizer, batches, steps=1000, log_every=100):
num_batches = batches.shape[0]
history = []
# Warmup: compile the training step
warmup_metrics = train_step(model, optimizer, batches[0])
block_tree(warmup_metrics)
start = time.perf_counter()
for step in range(steps):
tokens = batches[step % num_batches]
metrics = train_step(model, optimizer, tokens)
if step % log_every == 0 or step == steps - 1:
metrics = block_tree(metrics)
history.append(
{
"step": step,
"loss": float(metrics["loss"]),
"perplexity": float(metrics["perplexity"]),
}
)
block_tree(metrics)
elapsed = time.perf_counter() - start
batch_size = int(batches.shape[1])
tokens_per_step = batch_size * (SEQ_LEN - 1)
tokens_per_sec = steps * tokens_per_step / elapsed
return history, elapsed, tokens_per_sec
این سلول فقط دو تابع را تعریف میکند، بنابراین هیچ خروجی تولید نمیکند. مرحله بعدی جایی است که آنها اجرا میشوند.
۶. آموزش روی یک پردازنده گرافیکی واحد
برای ایجاد یک خط مبنا، با یک پردازنده گرافیکی (GPU) شروع کنید. jax.device_put کل آرایه دستهای را به gpu_devices[0] پین میکند، بنابراین قبل از اینکه به مقایسه چند پردازنده گرافیکی برسید، هیچ چیز بین دستگاهها پخش نمیشود.
BENCHMARK_STEPS = 500
STEPS_1GPU = BENCHMARK_STEPS
batches_1gpu = make_batches(train_data, PER_DEVICE_BATCH)
single_device = gpu_devices[0]
batches_1gpu = jax.device_put(batches_1gpu, single_device)
model_1gpu = TinyTransformer(
VOCAB_SIZE,
D_MODEL,
NUM_HEADS,
FFN_DIM,
NUM_LAYERS,
MAX_SEQ_LEN,
rngs=nnx.Rngs(1),
)
optimizer_1gpu = nnx.Optimizer(
model_1gpu, optax.adamw(LR, weight_decay=WEIGHT_DECAY), wrt=nnx.Param
)
history_1gpu, elapsed_1gpu, tps_1gpu = train_loop(
model_1gpu, optimizer_1gpu, batches_1gpu, steps=STEPS_1GPU
)
show_table(
["Step", "Loss", "Perplexity"],
[(h["step"], f"{h['loss']:.3f}", f"{h['perplexity']:.1f}") for h in history_1gpu],
title=f"Single-GPU training \u2014 {tps_1gpu:,.0f} tokens/sec",
aligns=["right", "right", "right"],
)
اولین فراخوانی، مرحله را کامپایل میکند، به همین دلیل است که train_loop قبل از شروع ساعت، گرم میشود. وقتی اجرا تمام شد، باید جدولی با یک ردیف برای هر مرحله ثبت شده ببینید که نشاندهنده کاهش loss و perplexity با پیشرفت آموزش است. عنوان جدول، توکنهای اندازهگیری شده بر ثانیه را برای این اجرا گزارش میدهد.
۷. یک مرحله را برای همه پردازندههای گرافیکی (GPU) یکسان در نظر بگیرید.
این یک الگوی موازیسازی دادهها است که در آن یک مش ایجاد میکنید، مدل را تکثیر میکنید و دادهها را در امتداد بُعد دستهای تقسیم میکنید. کد مرحله آموزش به هیچ وجه تغییر نمیکند و @nnx.jit موازیسازی را بر اساس نحوه قرارگیری آرایهها مدیریت میکند.
برای حفظ مقایسهی منصفانهی توان عملیاتی، اجراهای تکپردازندهای و چندپردازندهای از تعداد مراحل زمانبندیشدهی یکسانی استفاده میکنند. هر پردازندهی گرافیکی همچنان توالیهای PER_DEVICE_BATCH را در هر مرحله پردازش میکند، بنابراین اجرای چندپردازندهای، یک دستهی سراسری بزرگتر را پردازش میکند.
تکرار یک مدل NNX سه فراخوانی نیاز دارد. nnx.state وضعیت ماژول را به عنوان یک PyTree استخراج میکند، jax.device_put آن PyTree را در هر دستگاهی که دارای شاردینگ تکرار شده است قرار میدهد و nnx.update آن را دوباره در ماژول مینویسد. وضعیت بهینهساز نیز همین روند را طی میکند.
STEPS_MULTI = BENCHMARK_STEPS
GLOBAL_BATCH = PER_DEVICE_BATCH * NUM_DEVICES
mesh = Mesh(np.array(gpu_devices), ("data",))
replicated = NamedSharding(mesh, P())
data_sharding = NamedSharding(mesh, P(None, "data", None))
batches_multi = make_batches(train_data, GLOBAL_BATCH)
batches_multi = jax.device_put(batches_multi, data_sharding)
model_multi = TinyTransformer(
VOCAB_SIZE,
D_MODEL,
NUM_HEADS,
FFN_DIM,
NUM_LAYERS,
MAX_SEQ_LEN,
rngs=nnx.Rngs(1),
)
optimizer_multi = nnx.Optimizer(
model_multi, optax.adamw(LR, weight_decay=WEIGHT_DECAY), wrt=nnx.Param
)
# Replicate model and optimizer state across all GPUs
model_state = nnx.state(model_multi)
nnx.update(model_multi, jax.device_put(model_state, replicated))
opt_state = nnx.state(optimizer_multi)
nnx.update(optimizer_multi, jax.device_put(opt_state, replicated))
print(
f"Global batch: {GLOBAL_BATCH} ({PER_DEVICE_BATCH} per GPU \u00d7 {NUM_DEVICES} GPUs)"
)
print(f"Training batches: {batches_multi.shape}")
print()
history_multi, elapsed_multi, tps_multi = train_loop(
model_multi, optimizer_multi, batches_multi, steps=STEPS_MULTI
)
show_table(
["Step", "Loss", "Perplexity"],
[(h["step"], f"{h['loss']:.3f}", f"{h['perplexity']:.1f}") for h in history_multi],
title=f"Multi-GPU training \u2014 {tps_multi:,.0f} tokens/sec",
aligns=["right", "right", "right"],
)
شما باید خط دستهای سراسری را ببینید که توالیهای PER_DEVICE_BATCH را به ازای هر GPU ضربدر تعداد GPUها، شکل آرایه دستهای خرد شده و سپس یک جدول تلفات دوم با ستونهای مشابه اجرای تک GPU گزارش میدهد. عنوان آن توکنهای چند GPU/ثانیه را گزارش میدهد.
۸. توان عملیاتی را مقایسه کرده و منحنیها را رسم کنید
مدل و مرحله آموزش در هر دو اجرا یکسان بودند. فقط جایگذاری دادهها تغییر کرد. هر دو اجرا از تعداد مراحل زمانبندیشده یکسانی استفاده کردند و هر پردازنده گرافیکی همچنان توالیهای PER_DEVICE_BATCH را در هر مرحله پردازش میکرد، بنابراین اجرای چند پردازنده گرافیکی دسته کلی بزرگتری دارد.
ms_per_step_1gpu = elapsed_1gpu / STEPS_1GPU * 1e3
ms_per_step_multi = elapsed_multi / STEPS_MULTI * 1e3
speedup = tps_multi / tps_1gpu
show_table(
["", "1 GPU", f"{NUM_DEVICES} GPUs", "Ratio"],
[
("Batch size", PER_DEVICE_BATCH, GLOBAL_BATCH, f"{NUM_DEVICES}×"),
("Per-GPU batch", PER_DEVICE_BATCH, PER_DEVICE_BATCH, "same"),
("Timed steps", STEPS_1GPU, STEPS_MULTI, "same"),
("ms/step", f"{ms_per_step_1gpu:.2f}", f"{ms_per_step_multi:.2f}", f"{ms_per_step_1gpu / ms_per_step_multi:.2f}×"),
("Tokens/sec", f"{tps_1gpu:,.0f}", f"{tps_multi:,.0f}", f"{speedup:.2f}×"),
],
title="Throughput comparison",
aligns=["left", "right", "right", "right"],
)
if speedup > NUM_DEVICES * 1.25:
print(
f"Note: the measured speedup is superlinear (> {NUM_DEVICES}x). "
"For this small benchmark, treat that as a measurement artifact rather "
"than a general hardware-scaling claim."
)
show_bars(
[("1 GPU", tps_1gpu), (f"{NUM_DEVICES} GPUs", tps_multi)],
"Training throughput",
"tokens/s",
)
اجرای چند پردازنده گرافیکی (multi-GPU) بر حسب توکن بر ثانیه سریعتر است زیرا یک دسته سراسری بزرگتر را پردازش میکند در حالی که دسته هر پردازنده گرافیکی ثابت میماند. اگر افزایش سرعت اندازهگیری شده بیشتر از تعداد پردازندههای گرافیکی باشد، آن را به عنوان یک مصنوع معیار از تفاوتهای کامپایلر، طرحبندی و هسته در نظر بگیرید، نه یک تضمین مقیاسبندی کلی.
در مرحله بعد، پیشرفت آموزش را برای دو اجرا، همزمان با مشاهده توکنهای بیشتر توسط مدل، رسم کنید. محور x، توکنهای پردازششده را نشان میدهد نه مراحل آموزش خام، زیرا اجرای چند پردازنده گرافیکی از یک دسته کلی بزرگتر استفاده میکند و بنابراین دادههای بیشتری را در هر مرحله مشاهده میکند. هرچه میزان اتلاف و پیچیدگی کمتر باشد، بهتر است، بنابراین منحنیها نشان میدهند که هر تنظیم با توجه به میزان متن پردازششده، چقدر سریع بهبود مییابد.
fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(12, 4))
for label, history, batch_size in [
("1 GPU", history_1gpu, PER_DEVICE_BATCH),
(f"{NUM_DEVICES} GPUs", history_multi, GLOBAL_BATCH),
]:
tokens_m = [
(h["step"] + 1) * batch_size * (SEQ_LEN - 1) / 1e6 for h in history
]
losses = [h["loss"] for h in history]
perps = [h["perplexity"] for h in history]
ax1.plot(tokens_m, losses, "o-", label=label, markersize=4)
ax2.plot(tokens_m, perps, "o-", label=label, markersize=4)
ax1.set_xlabel("Tokens processed (millions)")
ax1.set_ylabel("Loss")
ax1.set_title("Training loss")
ax1.legend()
ax1.grid(True, alpha=0.25)
ax2.set_xlabel("Tokens processed (millions)")
ax2.set_ylabel("Perplexity")
ax2.set_title("Training perplexity")
ax2.legend()
ax2.grid(True, alpha=0.25)
fig.suptitle("Training curves vs tokens processed")
fig.tight_layout()
plt.show()
شما باید دو پنل را در کنار هم ببینید که در سمت چپ آن loss و در سمت راست آن perplexity قرار دارد. هر کدام با یک منحنی در هر اجرا، هر دو با افزایش توکنهای پردازش شده، روند نزولی دارند.
۹. ذخیره و بازیابی یک چک پوینت با Orbax
Orbax حالت مدل را به عنوان دایرکتوری از فایلهای آرایه ذخیره میکند. StandardCheckpointer سادهترین API آن است: یک فراخوانی برای ذخیره، یک فراخوانی برای بازیابی.
برای مدلهای NNX، پارامترها را با nnx.state(model, nnx.Param) استخراج میکنید، آن PyTree را ذخیره میکنید و بعداً آن را با nnx.update در یک مدل جدید بازیابی میکنید. ساختار گراف ( nnx.GraphDef ) ذخیره نمیشود و از تعریف کلاس پایتون میآید، بنابراین برای بازسازی مدل قبل از بارگذاری وزنها در آن، به TinyTransformer در محدوده نیاز دارید.
کد زیر کل چرخه حیات را اجرا میکند. این کد فقط پارامترهای مدل آموزشدیده را استخراج میکند، آنها را روی دیسک ذخیره میکند، یک درخت ShapeDtypeStruct میسازد که به Orbax میگوید چه شکلها و dtypeهایی را انتظار داشته باشد، در یک TinyTransformer که به تازگی مقداردهی اولیه شده است، بازیابی میکند و بررسی میکند که مدل بازیابی شده همان logits مدل اصلی را در یک ورودی آزمایشی کوچک تولید کند.
ckpt_dir = pathlib.Path("/tmp/jax-course/l7-checkpoints")
# Extract model parameters (not optimizer state)
model_params = nnx.state(model_multi, nnx.Param)
# Save
checkpointer = ocp.StandardCheckpointer()
if (ckpt_dir / "trained").exists():
import shutil
shutil.rmtree(ckpt_dir / "trained")
checkpointer.save(ckpt_dir / "trained", model_params)
print(f"Checkpoint saved to {ckpt_dir / 'trained'}")
# Create abstract target for restore
abstract_params = jax.tree.map(
lambda x: jax.ShapeDtypeStruct(x.shape, x.dtype),
model_params,
)
# Restore into a fresh model
model_restored = TinyTransformer(
VOCAB_SIZE,
D_MODEL,
NUM_HEADS,
FFN_DIM,
NUM_LAYERS,
MAX_SEQ_LEN,
rngs=nnx.Rngs(99),
)
restored_params = checkpointer.restore(ckpt_dir / "trained", abstract_params)
nnx.update(model_restored, restored_params)
# Test
test_input = jnp.zeros((1, 16), dtype=jnp.int32)
logits_original = model_multi(test_input)
logits_restored = model_restored(test_input)
max_diff = float(jnp.max(jnp.abs(logits_original - logits_restored)))
show_table(
["", "Value"],
[
("Checkpoint path", str(ckpt_dir / "trained")),
("Parameters saved", f"{sum(x.size for x in jax.tree.leaves(model_params)):,}"),
("Max |original \u2212 restored|", f"{max_diff:.2e}"),
("Match", "\u2713" if max_diff < 1e-5 else "\u2717"),
],
title="Orbax checkpoint save and restore",
)
شما باید مسیر نقطه بررسی، تعداد پارامترهای ذخیره شده، حداکثر اختلاف مطلق بین لاگیتهای اصلی و بازیابی شده و یک علامت تیک را در زمانی که این اختلاف کمتر از ۱e-۵ باشد، مشاهده کنید.
۱۰. متن شکسپیرگونه تولید کنید
مدل آموزشدیده، بایت بعدی را در هر موقعیت پیشبینی میکند. برای تولید متن، یک اعلان به آن میدهید، لوجیتها را در آخرین موقعیت میگیرید، یک توکن را نمونهبرداری میکنید، آن را اضافه میکنید و این کار را تکرار میکنید.
تابع generate ، ورودی خود را به MAX_SEQ_LEN منتقل میکند، بنابراین مسیر رو به جلوی کامپایل شده توسط JIT همیشه شکل ورودی یکسانی را بدون کامپایل مجدد با رشد دنباله میبیند. با توجه به توجه سببی، انتقال بعد از توکنهای واقعی، خروجی در موقعیتهای قبلی را تحت تأثیر قرار نمیدهد.
nnx.split ماژول را به یک graphdef و یک state PyTree تقسیم میکند تا بتوان آن را از طریق jax.jit ارسال کرد و nnx.merge ماژول را درون تابع کامپایل شده بازسازی میکند.
@jax.jit
def get_logits_jit(graphdef, model_state, tokens):
model = nnx.merge(graphdef, model_state)
return model(tokens)
def generate(model, prompt_text, max_new_tokens=300, temperature=0.8):
graphdef, model_state = nnx.split(model)
tokens = list(prompt_text.encode("utf-8"))
key = jax.random.key(42)
for _ in range(max_new_tokens):
context = tokens[-MAX_SEQ_LEN:]
padded = context + [0] * (MAX_SEQ_LEN - len(context))
input_arr = jnp.array([padded], dtype=jnp.int32)
logits = get_logits_jit(graphdef, model_state, input_arr)
next_logit = logits[0, len(context) - 1]
if temperature <= 0:
next_token = int(jnp.argmax(next_logit))
else:
key, subkey = jax.random.split(key)
next_token = int(jax.random.categorical(subkey, next_logit / temperature))
tokens.append(next_token)
return bytes(tokens).decode("utf-8", errors="replace")
# Put model on a single device for generation
gen_model = TinyTransformer(
VOCAB_SIZE,
D_MODEL,
NUM_HEADS,
FFN_DIM,
NUM_LAYERS,
MAX_SEQ_LEN,
rngs=nnx.Rngs(99),
)
nnx.update(gen_model, checkpointer.restore(ckpt_dir / "trained", abstract_params))
print("=== Prompt: 'ROMEO:' | temperature=0.8 ===")
print()
print(generate(gen_model, "ROMEO:", max_new_tokens=300, temperature=0.8))
print()
print("=== Prompt: 'To be, or not' | temperature=0.6 ===")
print()
print(generate(gen_model, "To be, or not", max_new_tokens=300, temperature=0.6))
شما باید دو بلوک متن تولید شده را ببینید، یکی برای هر اعلان.
انتظار متنی در سطح آثار شکسپیر را نداشته باشید. اگر مدل را در نظر بگیرید، خواهید دید که خطها تقریباً در جای درست قرار گرفتهاند، نام گوینده با حروف بزرگ نوشته شده، حروف انگلیسی زیادی نوشته شده و معنی بسیار کمی دارد.
۱۱. تمیز کردن
حجم کاری 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 حذف کنید.
۱۲. تبریک
شما یک مبدل رمزگشا با Flax NNX ساختید، آن را در هر دو GPU روی گره خود آموزش دادید، آن را با Orbax بررسی کردید و متن را از وزنهای بازیابی شده تولید کردید.
آنچه آموختهاید
- چگونه Flax NNX یک مدل را در ماژولهای قابل استفاده مجدد سازماندهی میکند:
nnx.Embed،nnx.Linear،nnx.LayerNormوnnx.MultiHeadAttention، ایجاد پارامتر و مسیر رو به جلو را مدیریت میکنند و مرحله آموزشnnx.value_and_gradبه همراهnnx.Optimizerاستفاده میکند. - چگونه توجه سببی با
is_causal=Trueموقعیتهای آینده را میپوشاند تا مدل فقط بتواند به عقب نگاه کند، و چرا قلابattention_fnبه یک پوشش**_نیاز دارد - چگونه آموزش موازی دادهها، مدل را تکرار میکند و دسته را دقیقاً مانند codelab 6، بدون تغییر مرحله آموزش، در GPUها تقسیم میکند
- نحوه ذخیره و بازیابی پارامترهای مدل توسط Orbax ، با استفاده از
nnx.stateوnnx.updateبه عنوان پل ارتباطی بین ماژولهای NNX و PyTrees ساده - چگونه میتوان توان عملیاتی را بر حسب توکن بر ثانیه، واحد طبیعی مدلهای زبانی، اندازهگیری کرد، در حالی که warmup و
block_until_readyهمچنان کار حفظ اعداد صحیح را انجام میدهند؟ - چگونه تولید متن، پیشبینیهای مدل را به صورت تک تک توکنها تغذیه میکند، و چگونه دما تصادفی بودن آن نمونهبرداری را کنترل میکند
مراحل بعدی
- Codelab 8: یک مدل JAX آموزش دیده را صادر و ارائه دهید که در آن این نقطه بازرسی را میگیرید و آن را برای ارائه با استنتاج JIT، کامپایل AOT و صادرات به فرمتهای قابل حمل آماده میکنید.
-
D_MODELتغییر دهید (128 یا 512 را امتحان کنید) وNUM_LAYERS(2 یا 6 را امتحان کنید) تغییر دهید، سپس دوباره اجرا کنید. ظرفیت بیشتر به معنای مراحل کندتر است و لایههای بیشتر به معنای حافظه بیشتر است. - دمای تولید را تغییر دهید (0.0، 0.5 و 1.2 را امتحان کنید). صفر حریصانه است و مقادیر بالای 1 خروجی تصادفیتری تولید میکنند.
- تغییر backend توجه: بدنه
causal_sdpaرا به فراخوانیjax.nn.dot_product_attention(..., is_causal=True, implementation="cudnn")تغییر دهید و از فعالسازیهای سازگار با bf16 استفاده کنید، سپس نحوه رفتار توجه ادغامشده cuDNN را مشاهده کنید.