GKE ব্যবহার করে NVIDIA GPU-তে আপনার প্রথম JAX প্রোগ্রামটি চালান।

১. ভূমিকা

জিপিইউ-তে জ্যাক্স শেখার পথ। ল্যাব ১: জিপিইউ-তে জ্যাক্স দিয়ে শুরু করা।

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:// :8884 খুলুন http:// :8884 আপনার ব্রাউজারে http:// :8884 টাইপ করুন এবং অনুরোধ করা হলে টোকেনটি পেস্ট করুন।

একটি নোটবুক তৈরি করুন

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

লক্ষ্য করার মতো তিনটি বিষয়:

  1. np এবং jnp এর ব্যবহার ছাড়া কোডটি হুবহু একই।
  2. y.device একটি CUDA ডিভাইস নির্দেশ করে — JAX অ্যারেটিকে স্বয়ংক্রিয়ভাবে GPU-তে স্থাপন করেছে, কারণ সেটিই ডিফল্ট ব্যাকএন্ড।
  3. 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-তে স্কেল করুন।

রেফারেন্স নথি