১. ভূমিকা

JAX হলো উচ্চ-ক্ষমতাসম্পন্ন নিউমেরিক্যাল কম্পিউটিংয়ের জন্য একটি পাইথন লাইব্রেরি। বাহ্যিকভাবে এটি NumPy-এর মতো মনে হলেও, এর আড়ালে এটি আপনার পাইথন ফাংশনগুলোকে ট্রেস করে, XLA দিয়ে সেগুলোকে কম্পাইল করে এবং ফলাফলটি NVIDIA GPU-এর মতো অ্যাক্সিলারেটরে চালায়।
এই কোডল্যাবে আপনি টেরাফর্ম ব্যবহার করে এনভিডিয়া এল৪ জিপিইউ সহ একটি গুগল কুবারনেটিস ইঞ্জিন ক্লাস্টার প্রস্তুত করবেন, সেই জিপিইউ নোডে অফিসিয়াল এনভিডিয়া জেএএক্স কন্টেইনারের ভিতরে জুপিটারল্যাব চালাবেন এবং আপনার প্রথম জেএএক্স কম্পিউটেশন লিখবেন। এর শেষে আপনার কাছে একটি কার্যকরী পরিবেশ থাকবে, যার উপর ভিত্তি করে এই আট-পর্বের সিরিজের বাকি অংশগুলো তৈরি হবে।
আপনি যা করবেন
- Terraform ব্যবহার করে একটি 2× NVIDIA L4 GPU নোড পুল সহ একটি GKE Standard ক্লাস্টার প্রস্তুত করুন।
- NVIDIA JAX কন্টেইনার ইমেজ থেকে GPU নোডে JupyterLab স্থাপন করুন
-
nvidia-smiএবংjax.devices()ব্যবহার করে GPU-টি এন্ড-টু-এন্ড যাচাই করুন। -
jax.numpyব্যবহার করে JAX অ্যারে কোড লিখুন এবং ফলাফলটি GPU-তে আছে কিনা তা নিশ্চিত করুন। - JAX-কে সংজ্ঞায়িত করে এমন তিনটি রূপান্তর প্রয়োগ করুন:
jax.jit,jax.gradএবংjax.vmap
আপনার যা যা লাগবে
- বিলিং সক্ষম একটি গুগল ক্লাউড প্রজেক্ট, এবং ওয়ার্কশপ ক্রেডিট অথবা জিপিইউ ব্যবহারের জন্য একটি রিজার্ভেশন।
- আপনার নির্বাচিত অঞ্চলে কমপক্ষে ২টি এনভিডিয়া এল৪ জিপিইউ-এর জন্য কোটা ( জিপিইউ কোটা কীভাবে চেক করবেন )
- পাইথন এবং নামপাই সম্পর্কে প্রাথমিক ধারণা থাকতে হবে। কুডা (CUDA) বা কুবারনেটিস (Kubernetes) অভিজ্ঞতার প্রয়োজন নেই।
সম্পূর্ণ করতে আনুমানিক সময়: ৬০ মিনিট ।
২. শুরু করার আগে
আপনার প্রকল্প নির্বাচন করুন
গুগল ক্লাউড কনসোলে , বিলিং সক্ষম করা আছে এমন একটি প্রজেক্ট নির্বাচন করুন বা তৈরি করুন।
ওপেন ক্লাউড শেল
একটি ক্লাউড শেল সেশন শুরু করতে অ্যাক্টিভেট ক্লাউড শেল (কনসোলের উপরের ডানদিকে থাকা টার্মিনাল আইকন)-এ ক্লিক করুন, তারপর এটিকে আপনার প্রজেক্টে নির্দেশ করুন:
gcloud config set project <YOUR_PROJECT_ID>
এই ধাপের সবকিছু ক্লাউড শেল -এ চলে, যেখানে আগে থেকেই gcloud , kubectl , terraform এবং git ইনস্টল করা আছে।
প্রয়োজনীয় API গুলি সক্রিয় করুন
এই কোডল্যাবের প্রয়োজনীয় প্রতিটি এপিআই একটি কমান্ডেই সক্রিয় করুন:
gcloud services enable \
container.googleapis.com \
compute.googleapis.com \
iam.googleapis.com \
cloudresourcemanager.googleapis.com \
logging.googleapis.com \
monitoring.googleapis.com
আপনার জিপিইউ কোটা নিশ্চিত করুন
টেরাফর্ম কনফিগারেশনে দুটি এনভিডিয়া এল৪ জিপিইউ (NVIDIA L4 GPU) চাওয়া হয়েছে। আপনার কোটা আছে কিনা তা নিশ্চিত করুন:
gcloud compute regions describe us-central1 \
--format="value(quotas.filter(metric:NVIDIA_L4_GPUS).limit)"
আপনার মান 2 বা তার বেশি দেখা উচিত। যদি আপনি 0 দেখেন, তাহলে চালিয়ে যাওয়ার আগে কোটা বৃদ্ধির জন্য অনুরোধ করুন ।
ওয়ার্কশপ রিপোজিটরি ক্লোন করুন
টেরাফর্ম মডিউল এবং কুবারনেটিস ম্যানিফেস্ট ওয়ার্কশপ রিপোজিটরিতে থাকে:
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 নোড পুল রয়েছে। - PersistentVolumeClaim, JupyterLab Pod, এবং একটি LoadBalancer Service-এর জন্য
deploy/jupyter.yaml
৩. Terraform ব্যবহার করে GPU ক্লাস্টারটি প্রস্তুত করুন।
GKE-তে একটি GPU নোডের জন্য বেশ কিছু জিনিস একসাথে সংযুক্ত করার প্রয়োজন হয়: একটি VPC-নেটিভ ক্লাস্টার, অ্যাক্সিলারেটর সংযুক্ত একটি নোড পুল, এবং নোডটিতে NVIDIA ড্রাইভার ইনস্টল করা। Terraform মডিউলটি এই তিনটি কাজই করে, ফলে আপনাকে কনসোলে ক্লিক করে যেতে হয় না।
আপনার প্রজেক্ট কনফিগার করুন
উদাহরণ ভেরিয়েবল ফাইলটি কপি করে আপনার প্রজেক্টে যুক্ত করুন:
cd terraform
cp terraform.tfvars.example terraform.tfvars
terraform.tfvars সম্পাদনা করুন এবং project_id সেট করুন। বাকি সবকিছুর ডিফল্ট এই কোডল্যাবের সাথে মেলে:
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
আপনি কী তৈরি করছেন তা বুঝুন
আবেদন করার আগে, main.tf এ থাকা নোড পুল ডেফিনিশনটি দেখুন। এই অংশটিই একটি সাধারণ নোডকে 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 GPU থাকে, এবং g2-standard-48 এ থাকে ৪টি। দ্বিতীয়ত, gpu_driver_installation_config ই নোডটিকে ব্যবহারযোগ্য করে তোলে — GKE উপযুক্ত NVIDIA ড্রাইভার ইনস্টল করে দেয়, ফলে আপনার Pod-কে শুধু CUDA ইউজার-স্পেস লাইব্রেরিগুলো ইনস্টল করতে হয়।
আবেদন করুন
terraform init
terraform apply
পরিকল্পনাটি পর্যালোচনা করুন এবং yes টাইপ করুন। ক্লাস্টার তৈরি এবং নোড-পুল প্রোভিশনিং-এ প্রায় ১০ মিনিট সময় লাগে। এই মুহূর্তে সামনের বিষয়গুলো পড়ে নেওয়া ভালো।
এটি শেষ হলে, ক্লাস্টার ক্রেডেনশিয়ালগুলি সংগ্রহ করুন যাতে kubectl নতুন ক্লাস্টারের সাথে যোগাযোগ করতে পারে:
$(terraform output -raw get_credentials_command)
নোডটিতে জিপিইউ আছে কিনা যাচাই করুন
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
৪. জিপিইউ নোডে জুপিটারল্যাব স্থাপন করুন
আপনার এখন একটি GPU নোড আছে, কিন্তু তাতে কিছুই চলছে না। deploy/jupyter.yaml এর ম্যানিফেস্টটি এমন একটি Pod শিডিউল করে, যা উভয় GPU-এর জন্য অনুরোধ করে এবং অফিসিয়াল NVIDIA JAX কন্টেইনার থেকে JupyterLab চালু করে।
ম্যানিফেস্ট প্রয়োগ করুন
cd ..
kubectl apply -f deploy/jupyter.yaml
এটি তিনটি অবজেক্ট তৈরি করে:
-
jax-workspace-pvcহলো/workspaceএ মাউন্ট করা একটি ৫০ জিবি স্থায়ী ভলিউম, যার ফলে পড রিস্টার্টের পরেও আপনার নোটবুকগুলো সুরক্ষিত থাকে। -
jax-jupyter, যে পডটিnvcr.io/nvidia/jax:26.04-maxtext-py3চালায় এবংnvidia.com/gpu: "2"জন্য অনুরোধ করে। -
jax-jupyter-svc, একটি লোডব্যালেন্সার যা ৮৮৮৪ পোর্টে JupyterLab-কে উন্মুক্ত করে।
GPU অনুরোধটি হলো গুরুত্বপূর্ণ লাইন:
resources:
limits:
nvidia.com/gpu: "2"
memory: "48Gi"
cpu: "12"
nvidia.com/gpu হলো GKE ডিভাইস প্লাগইন দ্বারা বিজ্ঞাপিত একটি বর্ধিত রিসোর্স। Kubernetes এই Pod-টিকে শুধুমাত্র সেই নোডেই শিডিউল করে যা এটিকে সরবরাহ করতে পারে, আর এভাবেই Pod-টি আপনার GPU নোড পুলে যুক্ত হয়।
পডটি প্রস্তুত হওয়া পর্যন্ত অপেক্ষা করুন।
কন্টেইনার ইমেজটি বড় এবং পডটি চালু হওয়ার সময় জুপিটারল্যাবও পিপ-ইনস্টল করে, তাই প্রথম পুল করতে কয়েক মিনিট সময় লাগে:
kubectl get pod jax-jupyter -w
STATUS Running হওয়া পর্যন্ত অপেক্ষা করুন, তারপর Ctrl+C চাপুন।
JupyterLab URL এবং টোকেন পান
সার্ভিসটির এক্সটার্নাল আইপি সংগ্রহ করুন:
kubectl get svc jax-jupyter-svc -w
EXTERNAL-IP পরিবর্তিত হওয়া পর্যন্ত অপেক্ষা করুন কোনো অ্যাড্রেসে যেতে, Ctrl+C চাপুন।
JupyterLab পড লগে একটি এককালীন লগইন টোকেন প্রিন্ট করে:
kubectl logs jax-jupyter | grep -o 'token=[a-z0-9]*' | head -1
http:// খুলুন http:// আপনার ব্রাউজারে http:// টাইপ করুন এবং অনুরোধ করা হলে টোকেনটি পেস্ট করুন।
একটি নোটবুক তৈরি করুন
JupyterLab-এ, /workspace এ একটি নতুন Python 3 নোটবুক তৈরি করুন। এই কোডল্যাবের বাকি অংশের প্রতিটি কোড ব্লক সেই নোটবুকের একটি সেলে যাবে।
৫. যাচাই করুন যে JAX, GPU-টিকে দেখতে পাচ্ছে।
যেকোনো JAX কোড লেখার আগে, নিশ্চিত করুন যে কন্টেইনারের ভেতর থেকে হার্ডওয়্যারটি দেখা যাচ্ছে। এই ধাপে ব্যর্থ হলে, পরবর্তী কোনো কিছুই কাজ করবে না।
হার্ডওয়্যার পরীক্ষা করুন
নোটবুক থেকে nvidia-smi চালান:
!nvidia-smi
আপনি ড্রাইভার সংস্করণ এবং বর্তমান মেমরি ব্যবহার সহ দুটি L4 এন্ট্রি দেখতে পাবেন।
এখন কম্পিউট ক্যাপাবিলিটি জানতে চান, এটি একটি দুই-অঙ্কের সংখ্যা যা হার্ডওয়্যারের জেনারেশনকে শনাক্ত করে। পরবর্তী কোডল্যাবগুলো এমন সব ফিচার ব্যবহার করে যা এর উপর নির্ভরশীল: cuDNN ফিউজড অ্যাটেনশনের জন্য 8.0 বা তার নতুন সংস্করণ এবং FP8-এর জন্য 9.0 বা তার নতুন সংস্করণ প্রয়োজন।
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 হলো একটি অ্যাডা লাভলেস জিপিইউ, তাই আপনি SM 8.9 দেখতে পাবেন।
JAX জিপিইউ খুঁজে পেয়েছে কিনা তা পরীক্ষা করুন।
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 | কনভোলিউশন, জিইএমএম এবং কালেক্টিভ-এর জন্য এক্সএলএ এনভিডিয়া লাইব্রেরিগুলোকে কল করে। |
CUDA ড্রাইভার এবং রানটাইম | জিপিইউ-তে কার্নেল লোড করে এবং ডিভাইস মেমরি পরিচালনা করে। |
আপনাকে প্রায় কখনোই নিজে CUDA লিখতে হয় না, কিন্তু বেশিরভাগ ওয়ার্কলোডের ক্ষেত্রে এর ফলাফল হাতে লেখা কার্নেলের সাথে প্রতিযোগিতামূলক। কোডল্যাব ২ ট্রেসিং এবং কম্পাইলেশন লেয়ার উন্মোচন করে এবং কোডল্যাব ৩ দেখায় কীভাবে একটি প্রোফাইলারের মাধ্যমে এই সবকিছুর এক্সিকিউশন পর্যবেক্ষণ করা যায়।
৬. GPU-তে JAX অ্যারে কোড লিখুন
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-তে স্থাপন করেছে, কারণ সেটিই ডিফল্ট ব্যাকএন্ড। - JAX তার নিজস্ব অ্যারে টাইপ (
jax.Array) রিটার্ন করে, কোনো NumPy অ্যারে নয়।np.asarray(y)কল করলে GPU থেকে CPU-তে ডেটা ট্রান্সফার হয়।
তৃতীয় পয়েন্টটির দিকে খেয়াল রাখুন। GPU এবং হোস্টের মধ্যে প্রতিটি আদান-প্রদানে সময় লাগে, এবং একটি JAX অ্যারে প্রিন্ট করলে সিনক্রোনাইজেশন বাধ্যতামূলক হয়ে পড়ে, কারণ প্রদর্শন করার জন্য পাইথনকে ভ্যালুটি ফেচ করতে হয়। একটি ছোট উদাহরণের জন্য এটা ঠিক আছে; কিন্তু একটি টাইমড লুপের ভেতরে এটি একটি আসল সমস্যা। কার্যকরী নিয়মটি হলো: jnp দিয়ে অ্যারে তৈরি করুন, jnp দিয়েই সেগুলোর ওপর কাজ করুন, এবং শুধুমাত্র যখন আপনার ভ্যালুগুলো দেখার সত্যিই প্রয়োজন হবে, তখনই NumPy-তে রূপান্তর করুন। কোডল্যাব ৩ আপনাকে দেখাবে কীভাবে একটি প্রোফাইলে অনিচ্ছাকৃত ট্রান্সফার শনাক্ত করতে হয়।
৭. jit, grad, এবং vmap প্রয়োগ করুন।
জিপিইউ-তে নামপাই (NumPy) বেশ কাজের, কিন্তু এককভাবে এটি বিকল্পগুলোর চেয়ে খুব বড় কোনো অগ্রগতি নয়। জ্যাক্স (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 থেকে শুরু করেছেন। কোডল্যাব ৪ এর উপর ভিত্তি করে একটি সম্পূর্ণ ট্রেনিং লুপ তৈরি করে।
jax.vmap একটি ব্যাচ জুড়ে ভেক্টরাইজ করে।
জিপিইউ ব্যাচড কাজ পছন্দ করে। অনেকগুলো ইনপুটে একটি ফাংশন প্রয়োগ করার সহজ উপায় হলো পাইথনের for লুপ ব্যবহার করা, কিন্তু এতে কার্নেলগুলো এক এক করে চালু হয় এবং জিপিইউ-এর রিসোর্স কমে যায়। jax.vmap একটিমাত্র উদাহরণের জন্য লেখা ফাংশনকে এমন একটি সংস্করণে রূপান্তর করে যা কোনো লুপ বা ম্যানুয়াল রিসেপিং ছাড়াই ব্যাচ আকারে কাজ করে।
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 ব্যাচ করো না (ব্রডকাস্ট করো), বরং অ্যাক্সিস 0 বরাবর xs ব্যাচ করো। এর ফলাফল লুপের মতোই, কিন্তু এটি একটি একক ব্যাচড GPU অপারেশন হিসেবে ডিসপ্যাচ হয়।
তাদের রচনা করুন
আসল পরাশক্তি হলো গঠন। রূপান্তরগুলো স্তূপীকৃত হয়:
fast_batched_grad = jax.jit(
jax.vmap(jax.grad(loss), in_axes=(None, 0, 0))
)
এক লাইনের কোডেই আপনি একটি কম্পাইলড, ভেক্টরাইজড, ডিফারেনশিয়েটেড ফাংশন পেয়ে যাবেন, যা একটি ব্যাচের প্রতিটি (x, y) জোড়ার জন্য আলাদা আলাদা গ্রেডিয়েন্ট রিটার্ন করে — ব্যাচড ট্রেনিংয়ের জন্য আপনার যা যা প্রয়োজন, তার প্রায় সবই এতে রয়েছে। কোডল্যাব ৪-এ একটি বাস্তব ডেটাসেটের উপর ঠিক এই প্যাটার্নটিকেই কাজে লাগানো হয়েছে।
৮. পরিষ্কার করুন
লোডব্যালেন্সার এবং পার্সিস্টেন্ট ভলিউম সহ জুপিটার ওয়ার্কলোডটি মুছে ফেলুন:
kubectl delete -f deploy/jupyter.yaml
ক্লাস্টার, নোড পুল, ভিপিসি এবং সার্ভিস অ্যাকাউন্ট ধ্বংস করুন:
cd terraform
terraform destroy
নির্দেশিত হলে yes টাইপ করুন। টিয়ারডাউন করতে কয়েক মিনিট সময় লাগবে।
অবশেষে, নিশ্চিত করুন যেন কিছুই পিছনে ফেলে না থাকে:
gcloud container clusters list
gcloud compute instances list
এই প্রজেক্টের জন্য উভয়ই খালি থাকা উচিত। যদি আপনি শুধু এই কোডল্যাবের জন্য একটি প্রজেক্ট তৈরি করে থাকেন, তাহলে আপনি এর পরিবর্তে ক্লাউড কনসোল থেকে পুরো প্রজেক্টটি মুছে ফেলতে পারেন।
৯. অভিনন্দন
আপনি একেবারে শূন্য থেকে একটি GPU ক্লাস্টার তৈরি করেছেন এবং তাতে আপনার প্রথম JAX প্রোগ্রামটি চালিয়েছেন।
আপনি যা শিখেছেন
- টেরাফর্ম ব্যবহার করে কীভাবে একটি NVIDIA L4 GPU নোড পুল সহ একটি GKE Standard ক্লাস্টার প্রোভিশন করতে হয়, যার মধ্যে নোডটিকে ব্যবহারযোগ্য করে তোলার জন্য
gpu_driver_installation_configব্লকটিও অন্তর্ভুক্ত। -
nvidia.com/gpuএক্সটেন্ডেড রিসোর্স ব্যবহার করে কীভাবে একটি পডকে জিপিইউ নোডে শিডিউল করবেন - JAX-on-GPU স্ট্যাকটি যেভাবে একসাথে কাজ করে: JAX ট্রেস, XLA কম্পাইল, এবং cuDNN, cuBLAS ও CUDA রানটাইম এক্সিকিউট হয়
-
nvidia-smi, compute capability,jax.devices(), এবংjax.default_backend()ব্যবহার করে কীভাবে একটি GPU এনভায়রনমেন্ট ভেরিফাই করবেন -
jax.numpyসাথে NumPy-এর মিল ও অমিলগুলো: অপরিবর্তনীয় অ্যারে,.at[...]আপডেট, ডিফল্ট ৩২-বিট ডেটাটাইপ, এবং হোস্ট ট্রান্সফারের খরচ -
jax.jit,jax.gradএবংjax.vmapকীভাবে প্রয়োগ ও কম্পোজ করতে হয়
পরবর্তী পদক্ষেপ
- কোডল্যাব ২:
jax.jitদিয়ে JAX কম্পাইলেশন নিয়ন্ত্রণ করুন, যেখানে আপনি শিখবেন কেন প্রথম কলটি ধীর হয়, কী কারণে রিকম্পাইল হয়, এবং কীভাবে শেপ স্থিতিশীল রাখা যায়। - JAX ইম্পোর্ট করার আগে
JAX_PLATFORMS=cpuব্যবহার করে দেখুন, এতে CPU রান করতে বাধ্য হবেন, এবং উপরেরjax.jitটাইমিংগুলোর সাথে তুলনা করুন। -
deploy/jupyter.yamlফাইলেmachine_type = "g2-standard-48"ওgpu_count = 4সেট করে এবংnvidia.com/gpuসাথে মিলিয়ে নোড পুলটিকে ৪টি L4 GPU-তে স্কেল করুন।