۱. مقدمه

در آزمایشگاه کد «اجرای اولین برنامه JAX خود روی پردازندههای گرافیکی NVIDIA با GKE»، شما یک تابع را در jax.jit روی پردازنده گرافیکی (GPU) قرار دادید و مشاهده کردید که اولین فراخوانی بسیار طولانیتر از هر فراخوانی پس از آن طول میکشد. این تصادفی نیست: JAX تابع پایتون شما را با متغیرهای انتزاعی ردیابی میکند ، برنامه ضبط شده را به XLA تحویل میدهد و فایل اجرایی کامپایل شدهای را که روی پردازنده گرافیکی اجرا میشود، ذخیره میکند. در این آزمایشگاه کد، آن فرآیند را باز میکنید، یاد میگیرید که چه چیزی یک ورودی را در حافظه پنهان کامپایل قرار میدهد و دو موردی را که بیشترین زمان را برای کاربران JAX صرف میکنند، مانند کامپایلهای مجدد تصادفی و جریان کنترل پایتون روی مقادیر ردیابی شده، برطرف میکنید.
کاری که انجام خواهید داد
- ردیابی ساعت با قرار دادن یک
printپایتون درون یک تابع jitted اتفاق میافتد. - هزینه کامپایل را در مقابل هزینه اجرای ذخیره شده در حافظه پنهان (cache) روی پردازنده گرافیکی (GPU) اندازهگیری کنید.
- شناسایی آنچه متعلق به کلید حافظه پنهان کامپایل است و چه چیزی باعث کامپایل مجدد میشود
- جریان کنترل پایتون را روی مقادیر ردیابی شده با
jnp.whereوjax.lax.condجایگزین کنید - شکلها را با
jax.lax.scan، padding و masking وstatic_argnumsپایدار نگه دارید - بررسی کنید که JAX چه چیزی را با
jax.make_jaxprردیابی کرده است
آنچه نیاز دارید
- یک پروژه Google Cloud با قابلیت پرداخت صورتحساب، و اعتبار کارگاه یا رزرو شامل استفاده از GPU
- سهمیه حداقل ۲ پردازنده گرافیکی NVIDIA L4 در منطقه انتخابی شما ( نحوه بررسی سهمیه پردازنده گرافیکی )
- تکمیل Codelab 1: اجرای اولین برنامه JAX خود روی پردازندههای گرافیکی NVIDIA با GKE یا یک محیط معادل JAX GPU
زمان تخمینی برای تکمیل: ۵۰ دقیقه .
۲. قبل از شروع
پروژه خود را انتخاب کنید
در کنسول گوگل کلود ، یک پروژه با قابلیت پرداخت فعال انتخاب یا ایجاد کنید.
پوسته ابری را باز کنید
برای شروع یک جلسه Cloud Shell ، روی Activate Cloud Shell (آیکون ترمینال در سمت راست بالای کنسول) کلیک کنید، سپس آن را به پروژه خود هدایت کنید:
gcloud config set project <YOUR_PROJECT_ID>
این codelab در همان محیط Codelab 1 اجرا میشود: اولین برنامه JAX خود را روی GPUهای NVIDIA با GKE اجرا کنید . اگر کلاستر GKE و JupyterLab Pod شما هنوز در حال اجرا هستند، از مرحله راهاندازی و تأیید GPU صرف نظر کنید. در غیر این صورت، همین حالا محیط را آماده کنید.
فراهم کردن محیط 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 به سلولی از آن دفترچه یادداشت میرود.
GPU را تنظیم و تأیید کنید
JAX، NumPy و چند کتابخانه کمکی استاندارد را وارد کنید، سپس تأیید کنید که روی یک GPU هستید.
import time
from functools import partial
import jax
import jax.numpy as jnp
import numpy as np
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"Devices: {devices}")
assert gpu_devices, f"This lab assumes a GPU backend. Available devices: {devices}"
print(f"GPU devices: {gpu_devices}")
شما باید gpu به عنوان backend پیشفرض و حداقل یک CudaDevice در لیست دستگاهها ببینید. این codelab فقط به یک GPU نیاز دارد، بنابراین اگر گره بیشتر در معرض دید قرار دهد، اشکالی ندارد.
۳. تابع jax.jit را در حال ردیابی تماشا کنید
وقتی یک تابع JAX ساده و بدون jit را فراخوانی میکنید، هر عملیات از طریق پایتون اجرا میشود و همزمان با اجرا به GPU ارسال میشود. jax.jit این وضعیت را تغییر میدهد. به جای اجرای تابع شما با آرایههای واقعی، تابع را ردیابی میکند : JAX آن را یک بار با متغیرهای انتزاعی که فقط یک شکل و یک dtype دارند فراخوانی میکند و هر عملیات JAX را که روی آن متغیرهای تصادفی انجام میدهید در یک نمایش میانی به نام jaxpr ثبت میکند.
JAX، jaxpr را به StableHLO کاهش میدهد، برنامه را به XLA کاهش میدهد و XLA یک فایل اجرایی بهینه شده برای دستگاه هدف کامپایل میکند. XLA ممکن است عملیات را با هم ترکیب کند، اما یک تابع کامپایل شده همچنان میتواند به چندین هسته GPU کاهش یابد. از آن به بعد، فراخوانی تابع مستقیماً به فایل اجرایی ذخیره شده در حافظه پنهان پرش میکند.
بنابراین هر فراخوانی JIT سه مرحله دارد:
فاز | چه کاری انجام میدهد؟ | چه اتفاقی میافتد؟ |
ردیابی | محاسبه را ثبت کنید | پایتون یک بار اجرا میشود و JAX هر عملیاتی را که روی متغیرهای انتزاعی انجام میشود، در یک |
کامپایل | پایینتر از یک فایل اجرایی GPU | JAX مقدار |
اجرا | استفاده مجدد از فایل اجرایی ذخیره شده در حافظه پنهان | هر فراخوانی بعدی با شکلها و نوع دادههای منطبق، از ردیابی و کامپایل صرفنظر میکند و برنامهی ذخیرهشده در حافظهی پنهان را اجرا میکند. |
ردیابی (Trace) دلیل اجرای print پایتون درون یک تابع JIT فقط در اولین فراخوانی است. ردیابی و کامپایل با هم دلیل کند بودن اولین فراخوانی هستند. اجرا (Execute) دلیل سریع بودن هر فراخوانی پس از آن است.
ردیابی را در عمل ببینید
به خودتان ثابت کنید که بدنه تابع فقط یک بار به ازای هر امضای ورودی اجرا میشود. یک print در سطح پایتون درون تابع قرار دهید: در حین ردیابی اجرا میشود، اما بخشی از برنامه کامپایل شده GPU نیست ، بنابراین بعداً با همان شکل و dtype فراخوانی میشود و چیزی چاپ نمیکند.
@jax.jit
def f(x):
"""Jitted demo function that prints during tracing so we can see exactly when JAX retraces."""
# This print runs during tracing only not on every GPU execution.
print(f" tracing with shape={x.shape} dtype={x.dtype}")
return x ** 2 + 1
print("Call 1 (new shape):")
_ = f(jnp.arange(4, dtype=jnp.float32)).block_until_ready()
print("Call 2 (same shape):")
_ = f(jnp.arange(4, dtype=jnp.float32)).block_until_ready()
print("Call 3 (new shape):")
_ = f(jnp.arange(5, dtype=jnp.float32)).block_until_ready()
شما باید خروجی مشابه زیر را ببینید:
Call 1 (new shape): tracing with shape=(4,) dtype=float32 Call 2 (same shape): Call 3 (new shape): tracing with shape=(5,) dtype=float32
دستور print در فراخوانی ۱، اولین باری که JAX شکل (4,) را با dtype float32 میبیند، و در فراخوانی ۳، اولین باری که شکل (5,) را میبیند، اجرا میشود. در فراخوانی ۲، JAX یک فایل اجرایی کامپایل شده موجود را پیدا میکند و از ردیابی و کامپایل صرف نظر میکند.
۴. سنجش کامپایل در مقابل اجرای کش شده
هزینه امضای جدید واقعی است، و این همان جایی است که گزارشهای کندی JAX از آن ناشی میشود. اندازهگیری کنید که چه مقدار از اولین فراخوانی، کامپایل و چه مقدار اجرا میشود.
تابع زیر 20 تابع غیرخطی را به هم متصل میکند، به طوری که کامپایل کردن به طور قابل توجهی گرانتر از اجرا است.
def heavy(x):
"""20 chained nonlinearities so the first-call compilation is visibly more expensive than the cached execution."""
y = x
for _ in range(20):
y = jnp.sin(y) * jnp.cos(y) + jnp.tanh(y)
return y
heavy_jit = jax.jit(heavy)
x = jnp.arange(1_000_000, dtype=jnp.float32)
# Empty in-process cache so a re-run shows the first-call compile cost again.
jax.clear_caches()
t0 = time.perf_counter()
_ = heavy_jit(x).block_until_ready()
first_ms = (time.perf_counter() - t0) * 1000
t0 = time.perf_counter()
for _ in range(20):
_ = heavy_jit(x).block_until_ready()
cached_ms = (time.perf_counter() - t0) * 1000 / 20
print(f"First call (compile + execute): {first_ms:8.2f} ms")
print(f"Cached call (execute only): {cached_ms:8.2f} ms")
print(f"Compilation cost (approx): {first_ms - cached_ms:8.2f} ms")
شما باید زمان اولین فراخوانی را ببینید که بسیار بزرگتر از زمان فراخوانی ذخیره شده است. فاصله بین آنها تقریباً برابر با زمانی است که XLA صرف کامپایل کردن آن کرده است.
برای توابع کوچک، این فاصله زمانی چند ده میلیثانیه است؛ برای یک مرحله کامل آموزش ترانسفورماتور، این فاصله به راحتی میتواند چندین ثانیه باشد. خبر خوب این است که شما برای هر ترکیب شکل و نوع داده، یک بار هزینه پرداخت میکنید، نه یک بار در هر فراخوانی. بقیه این کد در مورد این است که بیشتر از حد لازم هزینه پرداخت نکنید.
۵. بفهمید چه چیزی باعث کامپایل مجدد میشود
JAX حافظه پنهان کامپایل را بر اساس امضای ساختاری ورودیها تنظیم میکند: شکلها، نوعهای داده و هر آرگومان علامتگذاری شده به عنوان استاتیک. اگر امضا با امضایی که JAX قبلاً دیده است مطابقت داشته باشد، فایل اجرایی ذخیره شده در حافظه پنهان اجرا میشود. اگر چیزی تغییر کند، JAX دوباره ردیابی و کامپایل میکند.
سه چیز معمولاً باعث کامپایل مجدد میشوند:
بخش کلید کش | چه چیزهایی تغییر میکند | اثر |
شکل | شکل متفاوت | |
نوع D | نوع داده متفاوت | |
آرگومان استاتیک | مقدار استاتیک متفاوت | مقدار هر آرگومان |
مقادیر ورودیهای آرایه معمولی اهمیتی ندارند . دو آرایه float32 (32, 128) با محتوای کاملاً متفاوت به یک فایل اجرایی کامپایل شده یکسان برخورد میکنند.
به یک کامپایل مجدد نگاه کنید. حلقه زیر یک تابع jitted با پنج آرایه را فراخوانی میکند، که سه تا از آنها شکلی هستند که JAX هنوز ندیده است.
@jax.jit
def f(x):
"""Simple jitted scalar function used to demonstrate one compile per new input shape (a new dtype would trigger the same recompile)."""
return jnp.sum(x ** 2)
# clear JAX's in-process compilation cache.
jax.clear_caches()
# Feed in a few different shapes and measure each call.
shapes = [(100,), (200,), (100,), (200,), (300,)]
for s in shapes:
x = jnp.ones(s, dtype=jnp.float32)
t0 = time.perf_counter()
_ = f(x).block_until_ready()
dt = (time.perf_counter() - t0) * 1000
print(f"shape={str(s):8s} {dt:7.2f} ms")
شما باید سه فراخوانی آهسته، یکی برای هر شکل جدید، و دو فراخوانی سریع، برای (100,) و (200,) تکراری ببینید.
بارهای کاری واقعی همیشه این کار را بهطور تصادفی انجام میدهند: توالیهای با طول متغیر، آخرین دسته در یک دوره، خروجی توکنسازی نامنظم. راهحل تقریباً در همه موارد این است که اجازه ندهید شکل تغییر کند .
۶. جایگزینی جریان کنترل پایتون روی مقادیر ردیابی شده
در طول ردیابی، ورودیهای تابع شما آرایههای عینی نیستند. آنها مقادیر انتزاعی با شکل و نوع دادهی شناختهشده هستند. هر ساختار پایتون که نیاز به مقایسهی عددی آن محتواها داشته باشد ( if ، while ، bool(x) ، int(x) )، ردیابی را مختل میکند.
این چیزی است که به نظر میرسد. این ReLU از اساس اشتباه طراحی شده است:
@jax.jit
def relu_bad(x):
"""ReLU using a Python `if` on a traced value with JIT errors out at trace time."""
if x > 0:
return x
return jnp.zeros_like(x)
try:
print(relu_bad(jnp.array(1.0)))
except Exception as e:
print(f"{type(e).__name__}: {str(e).splitlines()[0]}")
شما باید خطای TracerBoolConversionError ببینید. این پیام به if اشاره میکند: JAX نمیتواند تصمیم بگیرد که کدام شاخه را نگه دارد وقتی مقدار انتزاعی است.
انتخاب را به صورت داده با jnp.where بیان کنید
راه حل این است که انتخاب را به صورت data بیان کنیم، نه به صورت جریان کنترل پایتون. برای یک انتخاب عنصری کوچک مانند ReLU، jnp.where تمیزترین ابزار است. هر دو شاخه همیشه اجرا میشوند و گزاره به JAX میگوید که در هر موقعیت از کدام یک استفاده کند.
@jax.jit
def relu(x):
"""ReLU using `jnp.where` with both branches are computed so tracing works."""
return jnp.where(x > 0, x, 0.0)
print(relu(jnp.array([-1.0, -0.5, 0.0, 0.5, 1.0])))
این بار هیچ خطایی وجود ندارد. دو ورودی منفی و صفر به صورت 0. برمیگردند و 0.5 و 1.0 بدون تغییر عبور میکنند.
یک شاخه واقعی با jax.lax.cond انتخاب کنید
برای شاخههایی که چیزهای بسیار متفاوتی را محاسبه میکنند، که اجرای هر دو بیفایده خواهد بود، از jax.lax.cond استفاده کنید. هر دو تابع شاخه ردیابی میشوند، اما در زمان اجرا lax.cond یک شرط XLA را نشان میدهد، بنابراین معمولاً فقط شاخه انتخاب شده اجرا میشود. یک نکته: تحت vmap ، cond ممکن است به یک عملیات شبیه به select تبدیل شود نه یک شاخه واقعی.
@jax.jit
def soft_or_sharp(x, sharp):
"""Switch between hard ReLU and softplus inside the compiled graph via `lax.cond`, controlled by a traced bool."""
# `sharp` is a scalar bool and lax.cond compiles to a real if-then-else
return jax.lax.cond(
sharp,
lambda x: jnp.where(x > 0, x, 0.0),
lambda x: jax.nn.softplus(x),
x,
)
x = jnp.array([-1.0, 0.5, 2.0])
print(f"sharp=True: {soft_or_sharp(x, jnp.array(True))}")
print(f"sharp=False: {soft_or_sharp(x, jnp.array(False))}")
شما باید دو آرایه متفاوت را ببینید: خط sharp=True ورودی منفی را به صفر محدود میکند، و خط sharp=False مقادیر مثبت کوچک softplus را در همه جا برمیگرداند.
۷. حلقههای بلند را با lax.scan فشرده نگه دارید.
برای حلقهها روی دادههای ردیابیشده، از متغیرهای جریان کنترل ساختاریافته مانند jax.lax.while_loop ، jax.lax.fori_loop و jax.lax.scan استفاده کنید.
در واقع، یک حلقه for پایتون با یک کران ایستا درون jit معتبر است، اما JAX هنگام ردیابی، حلقه را از حالت فشرده خارج میکند. این بدان معناست که ۲۰۰ تکرار حلقه تقریباً به ۲۰۰ بلوک تکراری در برنامه کامپایل شده تبدیل میشود. lax.scan حلقه را به عنوان یک مقدار اولیه حلقه مانند نگه میدارد، که معمولاً برای حلقههای طولانی با طول ثابت بسیار سریعتر کامپایل میشود.
# python_for_loop compile time scales with NUM_STEPS while scan_loop compile time stays roughly constant. Try NUM_STEPS = 2000 to see the gap widen.
NUM_STEPS = 200
@jax.jit
def python_for_loop(x):
"""Python `for` loop inside jit."""
y = x
for _ in range(NUM_STEPS):
y = jnp.sin(y) + 0.01 * y
return y
@jax.jit
def scan_loop(x):
"""Same logic expressed with `lax.scan`."""
def body(y, _):
y = jnp.sin(y) + 0.01 * y
return y, None
y, _ = jax.lax.scan(body, x, xs=None, length=NUM_STEPS)
return y
x = jnp.ones((1024,), dtype=jnp.float32)
jax.clear_caches()
t0 = time.perf_counter()
_ = python_for_loop(x).block_until_ready()
python_for_ms = (time.perf_counter() - t0) * 1000
jax.clear_caches()
t0 = time.perf_counter()
_ = scan_loop(x).block_until_ready()
scan_ms = (time.perf_counter() - t0) * 1000
print(f"Python for loop first call: {python_for_ms:8.2f} ms")
print(f"lax.scan first call: {scan_ms:8.2f} ms")
هر دو عدد شامل کامپایل میشوند و هر دو تابع، بازگشت یکسانی را محاسبه میکنند. حلقه پایتونِ باز شده باید یک برنامه بسیار بزرگتر را کامپایل کند، بنابراین اولین فراخوانی آن کندتر از دو فراخوانی دیگر است.
تفاوت مهم در ساختار زمان کامپایل است: حلقه پایتون در طول ردیابی باز میشود، در حالی که lax.scan به یک حلقه اولیه کاهش مییابد. برای حلقههای کوتاه، یک حلقه for پایتون اغلب خوب است. برای حلقههای مشتقپذیر طولانی، lax.scan معمولاً پیشفرض بهتری است.
۸. تثبیت اشکال با استفاده از padding و آرگومانهای استاتیک
حجمهای کاری واقعی از نظر شکل متفاوت هستند. برای مثال، آخرین دسته در یک دوره زمانی (epoch) کوچکتر است، توالیها طولهای متفاوتی دارند. اگر اجازه دهید شکل به JAX نشت کند، هر یک از این موارد باعث کامپایل مجدد میشود. راهحل استاندارد این است که ورودیها را به یک شکل ثابت وصل کنید و موقعیتهای استفاده نشده را بپوشانید .
MAX_LEN = 16
@jax.jit
def masked_mean(x, mask):
"""Mean of `x` ignoring positions where `mask==0`."""
# Always called with shape (MAX_LEN,) - no recompile when actual length varies
return jnp.sum(x * mask) / jnp.maximum(jnp.sum(mask), 1.0)
def pad(seq):
"""Right-pad a variable-length list of floats to `MAX_LEN` and return the padded array plus a 0/1 mask."""
actual_len = len(seq)
if actual_len > MAX_LEN:
raise ValueError(f"sequence length {actual_len} exceeds MAX_LEN={MAX_LEN}")
pad_len = MAX_LEN - actual_len
x = jnp.concatenate([
jnp.asarray(seq, dtype=jnp.float32),
jnp.zeros(pad_len, dtype=jnp.float32),
])
mask = jnp.concatenate([
jnp.ones(actual_len, dtype=jnp.float32),
jnp.zeros(pad_len, dtype=jnp.float32),
])
return x, mask
# Several different sequence lengths, but a single compiled function handles them all
for seq in [[1.0, 2.0, 3.0], [10.0] * 8, [5.0, -2.0]]:
x, mask = pad(seq)
print(f"len={len(seq):2d} mean={masked_mean(x, mask):.3f}")
شما باید برای هر دنباله یک خط ببینید که هر کدام فقط میانگین مقادیر واقعی را گزارش میدهند - فاصلهگذاری، میانگین را به سمت صفر نمیکشد.
هر سه فراخوانی به یک فایل اجرایی کامپایل شده یکسان برخورد میکنند، زیرا شکل روی دستگاه همیشه (MAX_LEN,) است. فقط ماسک تغییر میکند. این الگو در هر مقیاسی، از یک میانگین ساده ۱۶ عنصری در اینجا گرفته تا ماسکهای توجه پدگذاری شده در آموزش ترانسفورماتور در مقیاس بزرگ، خود را نشان میدهد.
وقتی میخواهید دوباره کامپایل کنید: static_argnums
با استفاده از Padding از کامپایل مجددی که هرگز درخواست نکردهاید، جلوگیری میکنید، اما با استفاده static_argnums عمداً چنین درخواستی را مطرح میکنید.
گاهی اوقات یک پارامتر واقعاً یک ثابت سمت پایتون مانند تعداد لایه، پرچم دقت یا اندازه هسته است و شما میخواهید JAX مقدار آن را در برنامه کامپایل شده ذخیره کند. آن آرگومانها را با static_argnums یا static_argnames برای آرگومانهای کلمه کلیدی علامتگذاری کنید. JAX مقدار آن آرگومانها را در کلید حافظه پنهان (cache key) هش میکند، بنابراین هر مقدار مجزا، فایل اجرایی کامپایل شده خود را دریافت میکند.
@partial(jax.jit, static_argnums=0)
def power(n: int, x):
"""Repeated squaring with `n` is static so JAX unrolls the loop and compiles a fresh program per value of `n`."""
# `n` is a Python int and JAX bakes it into the trace and unrolls the loop
y = x
for _ in range(n):
y = y * y
return y
jax.clear_caches()
for n in (2, 3, 2): # n=2 reuses the cache the second time
t0 = time.perf_counter()
_ = power(n, jnp.arange(4, dtype=jnp.float32)).block_until_ready()
print(f"n={n}: {(time.perf_counter() - t0) * 1000:7.2f} ms")
باید ببینید که دو فراخوانی اول کامپایل میشوند و فراخوانی سوم که n=2 را تکرار میکند، به سرعت برمیگردد.
هر مقدار جدید n باعث کامپایل میشود، اما برای پیکربندیهای ثابت دقیقاً همان چیزی است که شما میخواهید: حلقه کاملاً از حالت فشرده خارج میشود و XLA میتواند هر عملیات را ببیند. مصالحه ساده است: مقداری که به طور مداوم تغییر میکند را در static_argnums قرار ندهید، در غیر این صورت در هر فراخوانی دوباره کامپایل خواهید شد.
۹. ردیابی را با jax.make_jaxpr بررسی کنید
وقتی چیزی متفاوت از آنچه انتظار دارید کامپایل میشود، jax.make_jaxpr به شما امکان میدهد ردپا را قبل از اینکه XLA به آن برسد، ببینید. یک jaxpr یک نمایش میانی کامپایلر در سطح JAX است: یک نمایش تایپی و تابعی از آنچه JAX قبل از پایین آوردن به StableHLO و سپس XLA مرحلهبندی کرده است.
این کد نهایی و بهینهشدهی GPU نیست، اما برای درک آنچه JAX ردیابی کرده بسیار مفید است.
def f(x):
return jnp.tanh(x) * jnp.sin(x) + jnp.log1p(x * x)
print(jax.make_jaxpr(f)(jnp.arange(4, dtype=jnp.float32)))
شما باید یک برنامهی تایپشدهی کوچک را ببینید: یک متغیر اولیه در هر خط - tanh ، sin ، ضربها، log1p و جمع نهایی - که هر کدام با نوع آرایهای که تولید میکنند، حاشیهنویسی شدهاند.
اگر تا به حال شک کردهاید که JAX به دلیل تغییر غیرمنتظره یک شکل یا نوع داده (dtype) در حال کامپایل مجدد است، مقایسه دو jaxprs از یک فراخوانی "سریع" و یک فراخوانی "کند" معمولاً مقصر را مشخص میکند. همین ترفند برای فهمیدن اینکه چرا یک تبدیل مانند grad یا vmap کار بیشتری از آنچه در نظر داشتید تولید میکند، کار میکند.
۱۰. تمیز کردن
حجم کاری 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.jit ساختید: JAX تابع پایتون شما را ردیابی میکند، محاسبه ردیابی شده را به ورودی کامپایلر تبدیل میکند و فایل اجرایی کامپایل شده را برای تطبیق امضاهای ورودی ذخیره میکند.
آنچه آموختهاید
- چگونه یک تابع پایتون را با
jax.jitردیابی کنیم، و چرا عوارض جانبی پایتون مانند اجرایprintدر طول ردیابی به جای هر اجرا - چگونه زمان کامپایل را از زمان اجرای کش شده با استفاده از زمانبندی ساده
block_until_ready()تشخیص دهیم؟ - کلیدهای حافظه نهان کامپایل: ساختار ورودی PyTree، شکلها، انواع داده و مقادیر آرگومان استاتیک
- چگونه با جایگزینی شاخهبندی پایتون با
jnp.whereیاjax.lax.condاز خطاهای جریان کنترل ردیابیشده جلوگیری کنیم، و چراjnp.whereمیتواند NaNها را به گرادیان نشت دهد - چگونه
lax.scanیک حلقه طولانی با طول ثابت را فشرده نگه میدارد به جای اینکه آن را در برنامه کامپایل شده باز کند - چگونه اشکال ورودی را با استفاده از padding و masking پایدار کنیم، و چگونه ثابتهای سمت پایتون را با استفاده از
static_argnumsدر trace قرار دهیم، زمانی که عمداً یک فایل اجرایی جداگانه میخواهیم. - چگونه میتوان با استفاده از
jax.make_jaxprردیابی سطح JAX را بررسی کرد، زمانی که یک تابع متفاوت از آنچه انتظار دارید کامپایل میشود
مراحل بعدی
- در Codelab 3: پروفایل و اشکالزدایی JAX روی GPU با XProf و Nsight Systems، یاد خواهید گرفت که ردیابی، کامپایل و اجرای هسته را در یک پروفایل واقعی مشاهده کنید.
- مقدار
NUM_STEPSاز ۲۰۰ به ۲۰۰۰ افزایش دهید و مقایسه حلقه را دوباره اجرا کنید تا فاصله بینlax.scanو حلقه پایتون باز شده بیشتر شود. - یک مقدار پیوسته متغیر در
static_argnumsقرار دهید: فراخوانیpowerباnگرفته شده از شمارندهای که هر فراخوانی را افزایش میدهد، و مشاهده کنید که هر فراخوانی دوباره کامپایل میشود