১. ভূমিকা

"Run your first JAX program on NVIDIA GPUs with GKE" কোডল্যাবে, আপনি GPU-তে jax.jit এর মধ্যে একটি ফাংশন র্যাপ করেছিলেন এবং দেখেছিলেন যে প্রথম কলটি তার পরবর্তী প্রতিটি কলের চেয়ে অনেক বেশি সময় নিচ্ছে। এটি কোনো দুর্ঘটনা নয়: JAX আপনার পাইথন ফাংশনটিকে অ্যাবস্ট্রাক্ট প্লেসহোল্ডার দিয়ে ট্রেস করে , রেকর্ড করা প্রোগ্রামটি XLA- কে হস্তান্তর করে এবং GPU-তে চালানো কম্পাইল করা এক্সিকিউটেবলটি ক্যাশে করে রাখে। এই কোডল্যাবে আপনি সেই প্রসেসটি খুলবেন, কম্পাইল ক্যাশে কী এন্ট্রি তৈরি করে তা শিখবেন এবং এমন দুটি জিনিসের সমাধান করবেন যা JAX ব্যবহারকারীদের সবচেয়ে বেশি সময় নষ্ট করে, যেমন—অনিচ্ছাকৃত রিকম্পাইল এবং ট্রেস করা ভ্যালুগুলোর উপর পাইথন কন্ট্রোল ফ্লো।
আপনি যা করবেন
- jitted ফাংশনের ভিতরে একটি Python
printরেখে ট্রেসিং হতে দেখুন। - জিপিইউ-তে ক্যাশড এক্সিকিউশন খরচের সাথে কম্পাইলেশন খরচ পরিমাপ করুন।
- কম্পাইল ক্যাশ কী-এর অন্তর্ভুক্ত বিষয়গুলো শনাক্ত করুন এবং কী কারণে পুনরায় কম্পাইল শুরু হয় তা জানুন।
- ট্রেস করা মানগুলিতে পাইথন কন্ট্রোল ফ্লো
jnp.whereএবংjax.lax.condদিয়ে প্রতিস্থাপন করুন। -
jax.lax.scan, padding, masking এবংstatic_argnumsব্যবহার করে আকার স্থিতিশীল রাখুন। -
jax.make_jaxprদিয়ে JAX যা ট্রেস করেছে তা পরীক্ষা করুন।
আপনার যা যা লাগবে
- বিলিং সক্ষম একটি গুগল ক্লাউড প্রজেক্ট, এবং ওয়ার্কশপ ক্রেডিট অথবা জিপিইউ ব্যবহারের জন্য একটি রিজার্ভেশন।
- আপনার নির্বাচিত অঞ্চলে কমপক্ষে ২টি এনভিডিয়া এল৪ জিপিইউ-এর জন্য কোটা ( জিপিইউ কোটা কীভাবে চেক করবেন )
- কোডল্যাব ১ সম্পন্ন হলে: GKE ব্যবহার করে NVIDIA GPU-তে, অথবা এর সমতুল্য কোনো JAX GPU পরিবেশে আপনার প্রথম JAX প্রোগ্রামটি চালান।
সম্পূর্ণ করতে আনুমানিক সময়: ৫০ মিনিট ।
২. শুরু করার আগে
আপনার প্রকল্প নির্বাচন করুন
গুগল ক্লাউড কনসোলে , বিলিং সক্ষম করা আছে এমন একটি প্রজেক্ট নির্বাচন করুন বা তৈরি করুন।
ওপেন ক্লাউড শেল
একটি ক্লাউড শেল সেশন শুরু করতে অ্যাক্টিভেট ক্লাউড শেল (কনসোলের উপরের ডানদিকে থাকা টার্মিনাল আইকন)-এ ক্লিক করুন, তারপর এটিকে আপনার প্রজেক্টে নির্দেশ করুন:
gcloud config set project <YOUR_PROJECT_ID>
এই কোডল্যাবটি কোডল্যাব ১: GKE দিয়ে NVIDIA GPU-তে আপনার প্রথম JAX প্রোগ্রাম চালান- এর মতো একই পরিবেশে চলে। যদি আপনার GKE ক্লাস্টার এবং JupyterLab Pod এখনও চালু থাকে, তাহলে সরাসরি 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 = " তারপর project_id = " প্রস্তুত করুন এবং JupyterLab স্থাপন করুন:
terraform init
terraform apply
$(terraform output -raw get_credentials_command)
cd ..
kubectl apply -f deploy/jupyter.yaml
terraform apply প্রায় ১০ মিনিট সময় লাগে। এটি শেষ হলে, Pod এবং LoadBalancer-এর জন্য অপেক্ষা করুন, তারপর Pod log থেকে ওয়ান-টাইম JupyterLab টোকেনটি পড়ুন:
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 এ একটি নতুন পাইথন ৩ নোটবুক তৈরি করুন। এই কোডল্যাবের প্রতিটি কোড ব্লক সেই নোটবুকের একটি সেলে যাবে।
জিপিইউ সেট আপ এবং যাচাই করুন
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 এবং ডিভাইস লিস্টে অন্তত একটি CudaDevice ) দেখা উচিত। এই কোডল্যাবটির জন্য শুধুমাত্র একটি জিপিইউ প্রয়োজন, তাই নোডটি একাধিক জিপিইউ সরবরাহ করলেও কোনো সমস্যা নেই।
৩. jax.jit-কে আপনার ফাংশন ট্রেস করতে দেখুন।
যখন আপনি একটি সাধারণ, নন-জিটেড JAX ফাংশন কল করেন, তখন প্রতিটি অপারেশন পাইথনের মাধ্যমে চলে এবং এক্সিকিউশনের সাথে সাথে GPU-তে ডিসপ্যাচ হয়। jax.jit এই বিষয়টি পরিবর্তন করে। এটি আপনার ফাংশনটিকে আসল অ্যারে দিয়ে চালানোর পরিবর্তে, ফাংশনটিকে ট্রেস করে : JAX এটিকে একবার অ্যাবস্ট্রাক্ট প্লেসহোল্ডার দিয়ে কল করে, যেগুলোতে শুধু একটি শেপ এবং একটি dtype থাকে, এবং সেই প্লেসহোল্ডারগুলোর উপর আপনার করা প্রতিটি JAX অপারেশনকে jaxpr নামক একটি অন্তর্বর্তী রিপ্রেজেন্টেশনে রেকর্ড করে রাখে।
JAX, jaxpr-কে StableHLO-তে লোয়ার করে, সেই লোয়ার করা প্রোগ্রামটি XLA- এর কাছে হস্তান্তর করে, এবং XLA টার্গেট ডিভাইসের জন্য একটি অপ্টিমাইজড এক্সিকিউটেবল কম্পাইল করে। XLA অপারেশনগুলোকে ফিউজ করতে পারে, কিন্তু একটি কম্পাইল করা ফাংশন একাধিক GPU কার্নেলেও লোয়ার হতে পারে। এরপর থেকে, ফাংশনটি কল করলে তা সরাসরি ক্যাশ করা এক্সিকিউটেবলে চলে যায়।
সুতরাং প্রতিটি JIT কলের তিনটি পর্যায় থাকে:
পর্যায় | এটা যা করে | কী ঘটে |
ট্রেস | গণনাটি লিপিবদ্ধ করুন। | পাইথন একবার চলে এবং JAX অ্যাবস্ট্রাক্ট প্লেসহোল্ডারগুলোর উপর করা প্রতিটি অপারেশনকে একটি |
সংকলন | জিপিইউ এক্সিকিউটেবলে নামিয়ে আনা | JAX |
কার্যকর করুন | ক্যাশে করা এক্সিকিউটেবলটি পুনরায় ব্যবহার করুন | একই আকার ও ডেটাটাইপ সহ পরবর্তী প্রতিটি কল ট্রেস এবং কম্পাইল এড়িয়ে যায় এবং ক্যাশ করা প্রোগ্রামটি চালায়। |
ট্রেস-এর কারণেই একটি JIT ফাংশনের ভিতরে পাইথনের print শুধুমাত্র প্রথম কলে কার্যকর হয়। ট্রেস এবং কম্পাইল একসাথে প্রথম কলটিকে ধীরগতির করে তোলে। এক্সিকিউট-এর কারণেই এর পরের প্রতিটি কল দ্রুত হয়।
ট্রেসিং বাস্তবে দেখুন
নিজেকে প্রমাণ করে দেখুন যে ফাংশন বডি প্রতিটি ইনপুট সিগনেচারের জন্য শুধুমাত্র একবারই রান করে। ফাংশনের ভিতরে একটি পাইথন-স্তরের print রাখুন: এটি ট্রেসিংয়ের সময় এক্সিকিউট হয়, কিন্তু এটি কম্পাইল করা GPU প্রোগ্রামের অংশ নয় , তাই পরবর্তীতে একই শেপ এবং ডেটাটাইপ দিয়ে কল করলে কিছুই প্রিন্ট হয় না।
@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
কল ১-এ প্রিন্টটি ফায়ার হয়, যখন JAX প্রথমবার float32 dtype-সহ (4,) আকৃতিটি দেখে, এবং কল ৩-এ, যখন এটি প্রথমবার (5,) আকৃতিটি দেখে। কল ২-এ, JAX একটি বিদ্যমান কম্পাইল করা এক্সিকিউটেবল খুঁজে পায় এবং ট্রেসিং ও কম্পাইলেশন উভয়ই এড়িয়ে যায়।
৪. ক্যাশড এক্সিকিউশনের সাথে কম্পাইলেশনের তুলনা করে পরিমাপ করুন।
নতুন সিগনেচার কস্ট একটি বাস্তব বিষয়, এবং JAX স্লো রিপোর্টগুলো এখান থেকেই আসে। পরিমাপ করুন প্রথম কলের কতটুকু কম্পাইলেশন এবং কতটুকু এক্সিকিউশন।
নিচের ফাংশনটি ২০টি নন-লিনিয়ারিটিকে এমনভাবে শৃঙ্খলিত করে যে, এক্সিকিউশনের তুলনায় কম্পাইলেশন দৃশ্যত বেশি ব্যয়বহুল হয়।
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 ইনপুটগুলোর স্ট্রাকচারাল সিগনেচারের উপর ভিত্তি করে কম্পাইল ক্যাশে কী (key) সেট করে: যেমন তাদের শেপ (shape), তাদের ডেটাটাইপ (dtypes), এবং স্ট্যাটিক (static) হিসেবে চিহ্নিত যেকোনো আর্গুমেন্ট। যদি সিগনেচারটি JAX-এর আগে দেখা কোনো সিগনেচারের সাথে মিলে যায়, তবে ক্যাশে থাকা এক্সিকিউটেবলটি রান করে। যদি কোনো পরিবর্তন হয়, JAX ট্রেস করে এবং পুনরায় কম্পাইল করে।
সাধারণত তিনটি কারণে রিকম্পাইল হয়:
ক্যাশে মূল অংশ | কী পরিবর্তন হয় | প্রভাব |
আকৃতি | বিভিন্ন আকৃতি | |
ডিটাইপ | বিভিন্ন ডিটাইপ | |
স্থির আর্গুমেন্ট | ভিন্ন স্থির মান | যেকোনো |
সাধারণ অ্যারে ইনপুটের মান কোনো বিষয় নয় । সম্পূর্ণ ভিন্ন বিষয়বস্তু সহ দুটি (32, 128) float32 অ্যারে একই কম্পাইল করা এক্সিকিউটেবলকে প্রভাবিত করে।
পুনরায় কম্পাইল হওয়ার প্রক্রিয়াটি দেখুন। নিচের লুপটি পাঁচটি অ্যারে সহ একটি জিটেড ফাংশনকে কল করে, যার মধ্যে তিনটির গঠন 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,) এর জন্য দুটি দ্রুতগতির কল দেখতে পাবেন।
বাস্তব ওয়ার্কলোডগুলো প্রায়শই অনিচ্ছাকৃতভাবে এমনটা করে থাকে: যেমন পরিবর্তনশীল দৈর্ঘ্যের সিকোয়েন্স, একটি ইপকের শেষ ব্যাচ, বা অগোছালো টোকেনাইজেশন আউটপুট। প্রায় সব ক্ষেত্রেই এর সমাধান হলো , শেপ বা আকৃতিকে পরিবর্তন হতে না দেওয়া ।
৬. ট্রেস করা মানগুলিতে পাইথন কন্ট্রোল ফ্লো প্রতিস্থাপন করুন
ট্রেসিং করার সময়, আপনার ফাংশনের ইনপুটগুলো কোনো সুনির্দিষ্ট অ্যারে নয়। এগুলো হলো একটি নির্দিষ্ট আকার এবং ডেটাটাইপ (dtype) সহ বিমূর্ত মান। পাইথনের এমন যেকোনো গঠন যা এই মানগুলোর সংখ্যাগত তুলনা করে ( 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 স্টেটমেন্টটিকে নির্দেশ করে: যখন মানটি abstract হয়, তখন JAX কোন শাখাটি রাখবে তা স্থির করতে পারে না।
jnp.where ব্যবহার করে পছন্দটিকে ডেটা হিসেবে প্রকাশ করুন।
এর সমাধান হলো পছন্দটিকে পাইথন কন্ট্রোল ফ্লো হিসেবে নয়, বরং ডেটা হিসেবে প্রকাশ করা। 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 লাইনটি সর্বত্র ছোট ধনাত্মক সফটপ্লাস মান ফেরত দেয়।
৭. lax.scan ব্যবহার করে দীর্ঘ লুপগুলিকে সংহত রাখুন।
ট্রেস করা ডেটার উপর লুপ চালানোর জন্য, jax.lax.while_loop , jax.lax.fori_loop , এবং jax.lax.scan এর মতো স্ট্রাকচার্ড কন্ট্রোল-ফ্লো প্রিমিটিভ ব্যবহার করুন।
প্রকৃতপক্ষে, jit ভিতরে একটি স্ট্যাটিক বাউন্ড সহ পাইথনের for লুপ বৈধ, কিন্তু ট্রেসিং করার সময় 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 সাধারণত একটি ভালো ডিফল্ট বিকল্প।
৮. প্যাডিং এবং স্ট্যাটিক আর্গুমেন্টের সাহায্যে আকার স্থিতিশীল করুন
বাস্তব ওয়ার্কলোডের গঠন বিভিন্ন রকম হয়। উদাহরণস্বরূপ, একটি ইপকের শেষ ব্যাচটি ছোট হয়, সিকোয়েন্সগুলোর দৈর্ঘ্য ভিন্ন ভিন্ন হয়। আপনি যদি এই গঠনের ভিন্নতা 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
প্যাডিং ব্যবহার করে আপনি অনাকাঙ্ক্ষিত রিকম্পাইল এড়াতে পারেন, কিন্তু static_argnums ব্যবহার করে আপনি ইচ্ছাকৃতভাবে তা চেয়ে নেন।
কখনও কখনও একটি প্যারামিটার প্রকৃতপক্ষে পাইথনের নিজস্ব কনস্ট্যান্ট হয়, যেমন লেয়ার সংখ্যা, প্রিসিশন ফ্ল্যাগ বা কার্নেল সাইজ, এবং আপনি চান JAX যেন এর মানটিকে কম্পাইল করা প্রোগ্রামের মধ্যে অন্তর্ভুক্ত করে দেয়। এই আর্গুমেন্টগুলোকে static_argnums দিয়ে চিহ্নিত করুন, অথবা কীওয়ার্ড আর্গুমেন্টের জন্য static_argnames করুন। JAX এই আর্গুমেন্টগুলোর মানকে ক্যাশ কী-তে হ্যাশ করে রাখে, ফলে প্রতিটি স্বতন্ত্র মানের জন্য নিজস্ব কম্পাইল করা এক্সিকিউটেবল তৈরি হয়।
@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 দিয়ে ট্রেসটি পরীক্ষা করুন।
যখন কোনো কিছু আপনার প্রত্যাশার চেয়ে ভিন্নভাবে কম্পাইল হয়, তখন XLA সেটিকে পরিবর্তন করার আগেই jax.make_jaxpr আপনাকে তার ট্রেস দেখতে দেয়। একটি 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 রিকম্পাইল করছে, তাহলে একটি 'ফাস্ট' এবং একটি 'স্লো' কল থেকে পাওয়া দুটি jaxpr-এর তুলনা করলে সাধারণত মূল কারণটি চিহ্নিত করা যায়। grad বা vmap মতো কোনো ট্রান্সফরমেশন কেন আপনার ইচ্ছার চেয়ে বেশি কাজ করছে, তা বের করার ক্ষেত্রেও একই কৌশল কাজ করে।
১০. পরিষ্কার করুন
লোডব্যালেন্সার এবং পার্সিস্টেন্ট ভলিউম সহ জুপিটার ওয়ার্কলোডটি মুছে ফেলুন:
kubectl delete -f deploy/jupyter.yaml
ক্লাস্টার, নোড পুল, ভিপিসি এবং সার্ভিস অ্যাকাউন্ট ধ্বংস করুন:
cd terraform
terraform destroy
নির্দেশিত হলে yes টাইপ করুন, তারপর নিশ্চিত করুন যে পিছনে কিছু ফেলে রাখা হয়নি:
gcloud container clusters list
gcloud compute instances list
এই প্রজেক্টের জন্য উভয়ই খালি থাকা উচিত। যদি আপনি শুধু এই সিরিজের জন্য একটি প্রজেক্ট তৈরি করে থাকেন, তাহলে আপনি এর পরিবর্তে ক্লাউড কনসোল থেকে পুরো প্রজেক্টটি মুছে ফেলতে পারেন।
১১. অভিনন্দন
আপনি jax.jit পেছনের মানসিক মডেলটি তৈরি করেছেন: JAX আপনার পাইথন ফাংশনটি ট্রেস করে, ট্রেস করা গণনাকে কম্পাইলার ইনপুটে পরিণত করে এবং মিলে যাওয়া ইনপুট সিগনেচারের জন্য কম্পাইল করা এক্সিকিউটেবলটি ক্যাশে করে রাখে।
আপনি যা শিখেছেন
-
jax.jitব্যবহার করে কীভাবে একটি পাইথন ফাংশন ট্রেস করবেন, এবং কেনprintমতো পাইথন সাইড ইফেক্টগুলো প্রতিটি এক্সিকিউশনের পরিবর্তে ট্রেসিং চলাকালীন রান করে। - সহজ
block_until_ready()টাইমিং ব্যবহার করে কম্পাইলেশন টাইম এবং ক্যাশড এক্সিকিউশন টাইমের মধ্যে পার্থক্য কীভাবে করা যায় - কম্পাইল ক্যাশে যে বিষয়গুলোর উপর কাজ করে: ইনপুট PyTree স্ট্রাকচার, শেপ, ডেটাটাইপ এবং স্ট্যাটিক আর্গুমেন্টের মান।
- পাইথন ব্রাঞ্চিং-এর পরিবর্তে
jnp.whereবাjax.lax.condব্যবহার করে কীভাবে ট্রেসড-ভ্যালু কন্ট্রোল-ফ্লো ত্রুটি এড়ানো যায়, এবং কেনjnp.whereএকটি গ্রেডিয়েন্টে NaNs লিক করতে পারে। - কীভাবে
lax.scanএকটি দীর্ঘ নির্দিষ্ট দৈর্ঘ্যের লুপকে কম্পাইল করা প্রোগ্রামে প্রসারিত না করে সংহত রাখে - প্যাডিং এবং মাস্কিং ব্যবহার করে কীভাবে ইনপুট শেপ স্থিতিশীল করা যায়, এবং যখন আপনি ইচ্ছাকৃতভাবে একটি পৃথক এক্সিকিউটেবল চান তখন
static_argnumsব্যবহার করে কীভাবে পাইথন-সাইড কনস্ট্যান্টগুলিকে ট্রেসের মধ্যে অন্তর্ভুক্ত করা যায়। - যখন কোনো ফাংশন আপনার প্রত্যাশার চেয়ে ভিন্নভাবে কম্পাইল হয়, তখন
jax.make_jaxprব্যবহার করে কীভাবে JAX-স্তরের ট্রেস পরীক্ষা করবেন
পরবর্তী পদক্ষেপ
- কোডল্যাব ৩: XProf এবং Nsight Systems ব্যবহার করে GPU-তে JAX প্রোফাইল ও ডিবাগ করা-তে আপনি একটি বাস্তব প্রোফাইলে ট্রেসিং, কম্পাইলেশন এবং কার্নেল এক্সিকিউশন কীভাবে প্রদর্শিত হয় তা দেখতে শিখবেন।
-
lax.scanএবং আনরোল করা পাইথন লুপের মধ্যে ব্যবধান বাড়ানোর জন্য,NUM_STEPSমান 200 থেকে বাড়িয়ে 2000 করুন এবং লুপ তুলনাটি পুনরায় চালান। -
static_argnumsএ একটি ক্রমাগত পরিবর্তনশীল মান রাখুন: প্রতিটি কলে বৃদ্ধি পাওয়া একটি কাউন্টার থেকেnনিয়েpowerকল করুন, এবং দেখুন প্রতিটি কলই কীভাবে রিকম্পাইল হয়।