۱. مقدمه

JAX یک کتابخانه پایتون برای محاسبات عددی با کارایی بالا است. در ظاهر شبیه NumPy به نظر میرسد، اما در باطن توابع پایتون شما را ردیابی میکند، آنها را با XLA کامپایل میکند و نتیجه را روی شتابدهندههایی مانند پردازندههای گرافیکی NVIDIA اجرا میکند.
در این آزمایشگاه کد، شما یک کلاستر موتور کوبرنتیز گوگل را با پردازندههای گرافیکی سطح ۴ انویدیا با استفاده از Terraform آماده میکنید، JupyterLab را درون کانتینر رسمی انویدیا JAX روی آن گره پردازنده گرافیکی اجرا میکنید و اولین محاسبه JAX خود را مینویسید. در پایان، شما یک محیط کاری خواهید داشت که بقیه این مجموعه هشت قسمتی بر اساس آن ساخته میشود.
کاری که انجام خواهید داد
- فراهم کردن یک کلاستر استاندارد GKE با ۲ گره NVIDIA L4 GPU با استفاده از Terraform
- JupyterLab را روی گره GPU از تصویر کانتینر NVIDIA JAX مستقر کنید
- با استفاده از
nvidia-smiوjax.devices()پردازنده گرافیکی (GPU) را از ابتدا تا انتها بررسی کنید. - کد آرایه JAX را با
jax.numpyبنویسید و تأیید کنید که نتیجه روی GPU نمایش داده میشود. - سه تبدیلی که JAX را تعریف میکنند اعمال کنید:
jax.jit،jax.gradوjax.vmap
آنچه نیاز دارید
- یک پروژه Google Cloud با قابلیت پرداخت صورتحساب، و اعتبار کارگاه یا رزرو شامل استفاده از GPU
- سهمیه حداقل ۲ پردازنده گرافیکی NVIDIA L4 در منطقه انتخابی شما ( نحوه بررسی سهمیه پردازنده گرافیکی )
- آشنایی اولیه با پایتون و NumPy. نیازی به تجربه CUDA یا Kubernetes نیست.
زمان تخمینی برای تکمیل: ۶۰ دقیقه .
۲. قبل از شروع
پروژه خود را انتخاب کنید
در کنسول گوگل کلود ، یک پروژه با قابلیت پرداخت فعال انتخاب یا ایجاد کنید.
پوسته ابری را باز کنید
برای شروع یک جلسه Cloud Shell ، روی Activate Cloud Shell (آیکون ترمینال در سمت راست بالای کنسول) کلیک کنید، سپس آن را به پروژه خود هدایت کنید:
gcloud config set project <YOUR_PROJECT_ID>
همه چیز در این مرحله در Cloud Shell اجرا میشود که از قبل gcloud ، kubectl ، terraform و git روی آن نصب شدهاند.
فعال کردن API های مورد نیاز
فعال کردن هر API مورد نیاز این codelab در یک دستور:
gcloud services enable \
container.googleapis.com \
compute.googleapis.com \
iam.googleapis.com \
cloudresourcemanager.googleapis.com \
logging.googleapis.com \
monitoring.googleapis.com
سهمیه GPU خود را تأیید کنید
پیکربندی Terraform دو پردازنده گرافیکی NVIDIA L4 را درخواست میکند. تأیید کنید که سهمیه دارید:
gcloud compute regions describe us-central1 \
--format="value(quotas.filter(metric:NVIDIA_L4_GPUS).limit)"
شما باید مقداری برابر با 2 یا بالاتر ببینید. اگر 0 مشاهده کردید، قبل از ادامه، درخواست افزایش سهمیه دهید .
مخزن کارگاه را کلون کنید
ماژول Terraform و مانیفست Kubernetes در مخزن workshop موجود هستند:
git clone https://github.com/Google-Cloud-AI/partner-ai-nvidia.git
cd partner-ai-nvidia/05-workshops/jax-on-gpu
دو دایرکتوری مورد نیاز شما عبارتند از:
-
terraform/که شامل کلاستر استاندارد GKE، VPC، حساب سرویس گره و مجموعه گره L4 GPU است. -
deploy/jupyter.yamlبرای PersistentVolumeClaim، JupyterLab Pod و یک سرویس LoadBalancer
۳. آمادهسازی خوشه پردازنده گرافیکی (GPU) با Terraform
یک گره GPU در GKE به چندین چیز نیاز دارد که به هم متصل باشند: یک خوشه بومی VPC، یک مجموعه گره با یک شتابدهنده متصل و درایور NVIDIA نصب شده روی گره. ماژول Terraform هر سه را انجام میدهد، بنابراین لازم نیست از طریق کنسول کلیک کنید.
پروژه خود را پیکربندی کنید
فایل متغیرهای مثال را کپی کرده و آن را به پروژه خود ارجاع دهید:
cd terraform
cp terraform.tfvars.example terraform.tfvars
terraform.tfvars را ویرایش کنید و project_id را تنظیم کنید. مقادیر پیشفرض برای سایر موارد با این codelab مطابقت دارد:
project_id = "<YOUR_PROJECT_ID>"
region = "us-central1"
zone = "us-central1-a"
cluster_name = "jax-gpu-cluster"
machine_type = "g2-standard-24"
gpu_type = "nvidia-l4"
gpu_count = 2
بفهم چی داری خلق میکنی
قبل از اعمال، به تعریف node pool در main.tf نگاهی بیندازید. این بخشی است که یک node معمولی را به یک node GPU تبدیل میکند:
resource "google_container_node_pool" "gpu" {
name = "gpu-pool"
location = var.zone
cluster = google_container_cluster.primary.name
node_count = 1
node_config {
machine_type = var.machine_type
guest_accelerator {
type = var.gpu_type # nvidia-l4
count = var.gpu_count # 2
gpu_driver_installation_config {
gpu_driver_version = "DEFAULT"
}
}
disk_size_gb = 100
disk_type = "pd-balanced"
# ...
}
}
دو نکته مهم است. اول اینکه machine_type و gpu_count باید با هم مطابقت داشته باشند: g2-standard-24 دقیقاً با ۲ پردازنده گرافیکی L4 و g2-standard-48 با ۴ پردازنده گرافیکی عرضه میشود. دوم اینکه gpu_driver_installation_config چیزی است که گره را قابل استفاده میکند - GKE درایور NVIDIA منطبق را نصب میکند، بنابراین Pod شما فقط باید کتابخانههای فضای کاربر CUDA را بیاورد.
اعمال کردن
terraform init
terraform apply
طرح را مرور کنید و yes را تایپ کنید. ایجاد خوشه و تأمین گره-استخر حدود 10 دقیقه طول میکشد. این زمان خوبی برای مطالعهی پیش رو است.
وقتی کار تمام شد، اعتبارنامههای کلاستر را دریافت کنید تا kubectl با کلاستر جدید ارتباط برقرار کند:
$(terraform output -raw get_credentials_command)
تأیید کنید که گره دارای پردازندههای گرافیکی (GPU) است
kubectl get nodes -o custom-columns=\
NAME:.metadata.name,GPU:.status.allocatable.nvidia\\.com/gpu
شما باید خروجی مشابه زیر را ببینید:
NAME GPU gke-jax-gpu-cluster-gpu-pool-3f21a0b4-k7wq 2
۴. JupyterLab را روی گره GPU مستقر کنید
اکنون یک گره GPU دارید، اما هیچ چیزی روی آن اجرا نمیشود. فایل manifest در deploy/jupyter.yaml یک Pod را زمانبندی میکند که هر دو GPU را درخواست میکند و JupyterLab را از کانتینر رسمی NVIDIA JAX اجرا میکند.
مانیفست را اعمال کنید
cd ..
kubectl apply -f deploy/jupyter.yaml
این سه شیء ایجاد میکند:
-
jax-workspace-pvc، یک درایو دائمی ۵۰ گیگابایتی که در/workspaceنصب شده است تا نوتبوکهای شما پس از راهاندازی مجدد پاد، همچنان پابرجا بمانند. -
jax-jupyter، Pod کهnvcr.io/nvidia/jax:26.04-maxtext-py3را اجرا میکند وnvidia.com/gpu: "2"را درخواست میکند. -
jax-jupyter-svc، یک LoadBalancer که JupyterLab را روی پورت ۸۸۸۴ در معرض دید قرار میدهد.
درخواست GPU خط مهم است:
resources:
limits:
nvidia.com/gpu: "2"
memory: "48Gi"
cpu: "12"
nvidia.com/gpu یک منبع توسعهیافته است که توسط افزونه دستگاه GKE تبلیغ میشود. Kubernetes فقط این Pod را روی گرهای زمانبندی میکند که بتواند آن را برآورده کند، و به این ترتیب Pod روی استخر گره GPU شما قرار میگیرد.
صبر کنید تا پاد آماده شود
ایمیج کانتینر بزرگ است و پاد (Pod) نیز در هنگام شروع، JupyterLab را با pip نصب میکند، بنابراین اولین pull چند دقیقه طول میکشد:
kubectl get pod jax-jupyter -w
صبر کنید تا STATUS Running باشد، سپس Ctrl+C را فشار دهید.
دریافت آدرس اینترنتی و توکن JupyterLab
دریافت IP خارجی سرویس:
kubectl get svc jax-jupyter-svc -w
صبر کنید تا EXTERNAL-IP از ... تغییر کند. به یک آدرس، سپس Ctrl+C را فشار دهید.
JupyterLab یک توکن ورود یکبارمصرف را در لاگ Pod چاپ میکند:
kubectl logs jax-jupyter | grep -o 'token=[a-z0-9]*' | head -1
http:// را باز کنید http:// در مرورگر خود وارد کنید و توکن را در صورت درخواست، جایگذاری کنید.
یک دفترچه یادداشت ایجاد کنید
در JupyterLab، یک دفترچه یادداشت پایتون ۳ جدید در /workspace ایجاد کنید. هر بلوک کد در ادامه این codelab به سلولی از آن دفترچه یادداشت میرود.
۵. تأیید کنید که JAX پردازنده گرافیکی (GPU) را میبیند
قبل از نوشتن هرگونه کد JAX، مطمئن شوید که سختافزار از داخل کانتینر قابل مشاهده است. اگر این مرحله با شکست مواجه شود، هیچ چیز در پایین دست کار نخواهد کرد.
سختافزار را بررسی کنید
nvidia-smi را از نوت بوک اجرا کنید:
!nvidia-smi
شما باید دو ورودی L4 با نسخه درایور و میزان استفاده فعلی از حافظه را ببینید.
حالا قابلیت محاسبه (compute capability) را بپرسید، یک عدد دو رقمی که نسل سختافزار را مشخص میکند. آزمایشگاههای کد بعدی از ویژگیهایی استفاده میکنند که به آن وابسته هستند: cuDNN fused attention به نسخه ۸.۰ یا جدیدتر نیاز دارد و FP8 به نسخه ۹.۰ یا جدیدتر نیاز دارد.
import subprocess
def get_compute_capability() -> tuple[int, int]:
"""Query the compute capability of the first visible GPU."""
out = subprocess.check_output(
["nvidia-smi", "--query-gpu=compute_cap", "--format=csv,noheader"],
text=True,
)
major, minor = out.strip().split("\n")[0].split(".")
return int(major), int(minor)
SM_MAJOR, SM_MINOR = get_compute_capability()
print(f"Detected compute capability: SM {SM_MAJOR}.{SM_MINOR}")
if SM_MAJOR < 7:
print("WARNING: this course assumes SM 7.0+ (Volta or newer).")
else:
print("GPU is compatible with this course.")
پردازنده گرافیکی L4 از Ada Lovelace است، بنابراین باید SM 8.9 ببینید.
بررسی کنید که JAX پردازنده گرافیکی (GPU) را پیدا کرده باشد
import jax
import jax.numpy as jnp
devices = jax.devices()
gpu_devices = [d for d in devices if d.platform == "gpu"]
print(f"JAX version: {jax.__version__}")
print(f"Default backend: {jax.default_backend()}")
print(f"Available devices: {devices}")
assert gpu_devices, f"No GPU backend found. Available devices: {devices}"
print(f"GPU devices: {gpu_devices}")
شما باید خروجی مشابه زیر را ببینید:
JAX version: 0.7.2 Default backend: gpu Available devices: [CudaDevice(id=0), CudaDevice(id=1)] GPU devices: [CudaDevice(id=0), CudaDevice(id=1)]
اولین import jax چند ثانیه طول میکشد زیرا JAX زمان اجرای CUDA را مقداردهی اولیه کرده و دستگاهها را بررسی میکند.
نحوه قرارگیری قطعات در کنار هم
خود JAX هرگز مستقیماً با GPU کار نمیکند. این برنامهای میسازد که محاسبات شما را توصیف میکند و آن را به صورت یک پشته ارائه میدهد:
لایه | نقش |
جکس | تابع پایتون شما را به یک نمایش میانی ردیابی میکند |
ایکس ال ای | آن IR را به کد GPU بهینه شده کامپایل میکند. |
cuDNN، cuBLAS، NCCL | کتابخانههای NVIDIA XLA به دنبال کانولوشنها، GEMMها و کالکتیوها هستند |
درایور و زمان اجرا CUDA | هستهها را روی پردازنده گرافیکی بارگذاری میکند و حافظه دستگاه را مدیریت میکند. |
شما تقریباً هرگز خودتان CUDA نمینویسید، اما برای اکثر بارهای کاری نتیجه با هستههای دستنویس قابل رقابت است. Codelab 2 لایه ردیابی و کامپایل را باز میکند و codelab 3 به شما نشان میدهد که چگونه تمام اجرای آن را در یک پروفایلر مشاهده کنید.
۶. نوشتن کد آرایه JAX روی GPU
سریعترین راه برای آشنایی با JAX این است که توجه کنید چه مقدار از آن فقط NumPy است. همان سازندهها و قوانین پخش اعمال میشوند. چیزی که تغییر میکند محل قرارگیری آرایه و نحوه اجرای محاسبات است.
مقایسه NumPy و JAX در کنار هم
import numpy as np
# NumPy: runs on the CPU, stored in host memory
x_np = np.arange(8, dtype=np.float32)
y_np = np.sin(x_np) ** 2 + np.cos(x_np) ** 2
print(f"NumPy result: {y_np}")
print(f"NumPy device: CPU (host memory)")
print()
# JAX: same code, different array library
x = jnp.arange(8, dtype=jnp.float32)
y = jnp.sin(x) ** 2 + jnp.cos(x) ** 2
print(f"JAX result: {y}")
print(f"JAX device: {y.device}")
# Sanity check: the two answers should agree
np.testing.assert_allclose(y_np, np.asarray(y), atol=1e-6)
print()
print("NumPy and JAX agree.")
شما باید خروجی مشابه زیر را ببینید:
JAX result: [1. 1. 1. 1. 1. 1. 1. 1.] JAX device: cuda:0
سه نکته که باید به آنها توجه کنید:
- کد به جز
npدر مقابلjnpیکسان است. -
y.deviceیک دستگاه CUDA را گزارش میدهد - JAX آرایه را به طور خودکار روی GPU قرار داده است زیرا این backend پیشفرض است. - JAX نوع آرایه خودش (
jax.Array) را برمیگرداند، نه یک آرایه NumPy. فراخوانیnp.asarray(y)باعث انتقال GPU به CPU میشود.
به نکته سوم توجه کنید. هر عبور بین GPU و میزبان زمان میبرد و چاپ یک آرایه JAX باعث همگامسازی میشود، زیرا پایتون برای نمایش مقدار باید آن را دریافت کند. برای یک مثال کوچک، خوب است؛ یک مشکل واقعی در یک حلقه زمانبندی شده. قانون کار این است: آرایهها را با jnp ایجاد کنید، با jnp روی آنها عملیات انجام دهید و فقط زمانی که واقعاً نیاز به بررسی مقادیر دارید، آنها را به NumPy تبدیل کنید. Codelab 3 به شما نشان میدهد که چگونه انتقالهای تصادفی را در یک پروفایل تشخیص دهید.
۷. اعمال jit، grad و vmap
NumPy روی GPU مفید است، اما به خودی خود جهش بزرگی نسبت به جایگزینهایش نیست. چیزی که JAX را متمایز میکند، تبدیلهای تابع آن است: عملگرهایی که یک تابع پایتون را میگیرند و یک تابع جدید با قدرتهای اضافی برمیگردانند. سه تا از آنها در هر آزمایشگاه کد باقی مانده وجود دارند.
jax.jit تابع شما را کامپایل میکند
وقتی یک تابع ساده JAX را فراخوانی میکنید، عملیاتها یکییکی به GPU ارسال میشوند. هر ارسال سربار دارد و عملیاتهای کوچک باعث میشوند GPU کمتر مورد استفاده قرار گیرد. jax.jit مدل اجرا را تغییر میدهد: JAX تابع شما را ردیابی میکند، XLA آن را به یک فایل اجرایی بهینه کامپایل میکند و JAX آن فایل اجرایی را در فراخوانیهای بعدی با شکلها و نوعهای داده سازگار دوباره استفاده میکند.
اولین فراخوانی به دلیل کامپایل شدن کند است. فراخوانیهای بعدی سریع هستند.
import time
def f(x):
"""Compose tanh, sin, and log1p so XLA has multiple ops to fuse when jitted."""
return jnp.tanh(x) * jnp.sin(x) + jnp.log1p(x * x)
x = jnp.arange(1_000_000, dtype=jnp.float32)
# Eager: one kernel launch per operation
_ = f(x).block_until_ready() # warm up
t0 = time.perf_counter()
for _ in range(10):
y = f(x).block_until_ready()
eager_ms = (time.perf_counter() - t0) * 1000 / 10
print(f"Eager: {eager_ms:6.3f} ms / call")
# Compiled: optimized executable, often with fused operations
f_jit = jax.jit(f)
_ = f_jit(x).block_until_ready() # first call compiles
t0 = time.perf_counter()
for _ in range(10):
y = f_jit(x).block_until_ready()
jit_ms = (time.perf_counter() - t0) * 1000 / 10
print(f"jax.jit (cached): {jit_ms:6.3f} ms / call")
print(f"Speedup: {eager_ms / jit_ms:6.1f}x")
افزایش سرعت دقیق به اندازه و شکل محاسبات شما بستگی دارد، اما الگوی کلی آن این است: JAX مشتاق راحت است، JAX کامپایل شده سریع است.
jax.grad به طور خودکار تمایز قائل میشود
آموزش یک شبکه عصبی به معنای محاسبه گرادیانهای یک زیان نسبت به پارامترها است. هر تابع اسکالر را به jax.grad ارسال کنید و این تابع یک تابع جدید را برمیگرداند که مشتق را محاسبه میکند.
def loss(w, x, y):
"""Mean squared error of `w*x` vs `y`; scalar loss for the `jax.grad` demo below."""
pred = w * x
return jnp.mean((pred - y) ** 2)
w = jnp.array(0.5)
xs = jnp.array([1.0, 2.0, 3.0, 4.0])
ys = jnp.array([2.0, 4.0, 6.0, 8.0])
# grad returns a function with the same signature, differentiating w.r.t. the first argument
dloss_dw = jax.grad(loss)
print(f"loss(w=0.5): {loss(w, xs, ys):.4f}")
print(f"dloss/dw: {dloss_dw(w, xs, ys):.4f}")
# Sanity check against a finite-difference approximation
eps = 1e-3
fd = (loss(w + eps, xs, ys) - loss(w - eps, xs, ys)) / (2 * eps)
print(f"finite diff: {fd:.4f} (should match)")
گرادیان منفی است، که به بهینهساز میگوید افزایش w باعث کاهش تلفات میشود - کاملاً درست است، زیرا رابطه واقعی y = 2x است و شما از w = 0.5 شروع کردهاید. Codelab 4 یک حلقه آموزشی کامل در این مورد ایجاد میکند.
jax.vmap به صورت برداری در یک دسته قرار میگیرد
پردازندههای گرافیکی (GPU) کار دستهای (batch-work) را دوست دارند. روش ساده برای اعمال یک تابع به ورودیهای زیاد، استفاده از حلقه for پایتون است، اما این حلقه هستهها را یکی یکی اجرا میکند و پردازنده گرافیکی (GPU) را از کار میاندازد. jax.vmap تابعی را که برای یک مثال نوشته شده است، میگیرد و نسخهای را برمیگرداند که روی یک دسته (batch) عمل میکند، بدون حلقه و بدون تغییر شکل دستی.
def predict(W, x):
"""Tanh of a single-example matrix-vector product; vmapped below to batch over many `x`."""
# Single example: W is (out, in), x is (in,) -> result is (out,)
return jnp.tanh(W @ x)
key_w, key_x = jax.random.split(jax.random.key(0))
W = jax.random.normal(key_w, (4, 3))
xs = jax.random.normal(key_x, (10, 3)) # batch of 10 examples
# Without vmap: a Python loop, one kernel launch per example
ys_loop = jnp.stack([predict(W, x) for x in xs])
# With vmap: batch over the leading axis of xs, share W across the batch
batched_predict = jax.vmap(predict, in_axes=(None, 0))
ys_vmap = batched_predict(W, xs)
print(f"ys_loop shape: {ys_loop.shape}")
print(f"ys_vmap shape: {ys_vmap.shape}")
np.testing.assert_allclose(np.asarray(ys_loop), np.asarray(ys_vmap), atol=1e-6)
print("vmap matches the explicit loop.")
آرگومان in_axes=(None, 0) میگوید: W را دستهای (پخش) نکن، xs را در امتداد محور 0 دستهای کن. نتیجه با حلقه یکسان است، اما به عنوان یک عملیات GPU دستهای واحد ارسال میشود.
آنها را بنویسید
قدرت فوقالعاده واقعی، ترکیببندی است. پشته تبدیلها:
fast_batched_grad = jax.jit(
jax.vmap(jax.grad(loss), in_axes=(None, 0, 0))
)
یک خط به شما یک تابع کامپایل شده، برداری و مشتقگیری شده میدهد که برای هر جفت (x, y) در یک دسته، یک گرادیان به ازای هر نمونه برمیگرداند - بیشتر چیزی که برای آموزش دستهای نیاز دارید. Codelab 4 دقیقاً همین الگو را روی یک مجموعه داده واقعی به کار میگیرد.
۸. تمیز کردن
حجم کاری 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 حذف کنید.
۹. تبریک
شما یک کلاستر GPU را از ابتدا آماده کردید و اولین برنامه JAX خود را روی آن اجرا کردید.
آنچه آموختهاید
- نحوه آمادهسازی یک کلاستر استاندارد GKE با یک استخر گره NVIDIA L4 GPU با استفاده از Terraform، شامل بلوک
gpu_driver_installation_configکه گره را قابل استفاده میکند - نحوه زمانبندی یک پاد روی یک گره GPU با منبع توسعهیافته
nvidia.com/gpu - نحوهی قرارگیری پشتهی JAX-on-GPU: ردیابیهای JAX، کامپایلهای XLA، و اجرای cuDNN، cuBLAS و CUDA در زمان اجرا
- نحوه تأیید محیط پردازنده گرافیکی (GPU) با
nvidia-smi، compute capability،jax.devices()وjax.default_backend() - کجا
jax.numpyبا NumPy مطابقت دارد و کجا متفاوت است: آرایههای تغییرناپذیر، بهروزرسانیهای.at[...]، انواع داده پیشفرض ۳۲ بیتی و هزینه انتقال میزبان - نحوه اعمال و ترکیب
jax.jit،jax.gradوjax.vmap
مراحل بعدی
- Codelab 2: کنترل کامپایل JAX با
jax.jitکه در آن یاد خواهید گرفت که چرا اولین فراخوانی کند است، چه چیزی باعث کامپایل مجدد میشود و چگونه شکلها را پایدار نگه دارید - قبل از وارد کردن JAX برای اجرای اجباری CPU،
JAX_PLATFORMS=cpuرا امتحان کنید و زمانبندیهایjax.jitبالا را مقایسه کنید. - با تنظیم
machine_type = "g2-standard-48"وgpu_count = 4و تطبیقnvidia.com/gpuدرdeploy/jupyter.yamlمجموعه گرهها را به ۴ پردازنده گرافیکی L4 مقیاسبندی کنید.