cuDNN এবং TransformerEngine ব্যবহার করে GPU-তে অ্যাটেনশনের গতি বাড়ান

১. ভূমিকা

জিপিইউ-তে জ্যাক্স শেখার পথ। ল্যাব ৫: জিপিইউ-তে অ্যাটেনশন।

এই কোডল্যাবে আপনি এনভিডিয়া জিপিইউ-তে অ্যাটেনশন গণনার উপর মনোযোগ দেবেন। আপনি একটি সাধারণ ইমপ্লিমেন্টেশন দিয়ে শুরু করবেন, সেটিকে JAX-এর বিল্ট-ইন ফিউজড অ্যাটেনশন দিয়ে প্রতিস্থাপন করবেন, এবং তারপর একই অপারেশনের জন্য এনভিডিয়ার cuDNN ব্যাকএন্ড সক্রিয় করবেন। শেষে আপনি জানতে পারবেন প্রতিটি স্তরে কী পরিবর্তন হয়, এবং সমস্যা বড় হওয়ার সাথে সাথে জিপিইউ-অপ্টিমাইজড পাথগুলো কতটা দ্রুততর হয়।

আপনি যা করবেন

  • মৌলিক JAX অপারেশন ব্যবহার করে শুরু থেকে স্কেলড ডট-প্রোডাক্ট অ্যাটেনশন বাস্তবায়ন করুন।
  • এটিকে JAX-এর অন্তর্নির্মিত ফিউজড কার্নেল jax.nn.dot_product_attention দিয়ে প্রতিস্থাপন করুন।
  • implementation="cudnn" দিয়ে cuDNN ব্যাকএন্ডকে বাধ্যতামূলক করুন এবং কজাল মাস্কিং যোগ করুন।
  • কোথায় ফিউজড কার্নেলগুলো লাভজনক হয় তা দেখতে সুইপ সিকোয়েন্সের দৈর্ঘ্য এবং ব্যাচ সাইজ পরীক্ষা করুন।
  • মাল্টি-হেড অ্যাটেনশন শেপ MHA, GQA, এবং MQA-এর তুলনা করুন
  • এনভিডিয়া ট্রান্সফরমারইঞ্জিনের মনোযোগকে বেঞ্চমার্ক করুন এবং এর FP8 কোড পাথ পরিদর্শন করুন।

আপনার যা যা লাগবে

  • বিলিং সক্ষম একটি গুগল ক্লাউড প্রজেক্ট, এবং ওয়ার্কশপ ক্রেডিট অথবা জিপিইউ ব্যবহারের জন্য একটি রিজার্ভেশন।
  • আপনার নির্বাচিত অঞ্চলে কমপক্ষে ২টি এনভিডিয়া এল৪ জিপিইউ-এর জন্য কোটা ( জিপিইউ কোটা কীভাবে চেক করবেন )
  • কোডল্যাব ১ থেকে ৪ সম্পন্ন করা, অথবা একটি সমতুল্য CUDA-সক্ষম JAX GPU পরিবেশ।
  • ম্যাট্রিক্স গুণন এবং সফটম্যাক্স সম্পর্কে ধারণা থাকতে হবে। CUDA বা cuDNN-এর কোনো অভিজ্ঞতার প্রয়োজন নেই।

সম্পূর্ণ করতে আনুমানিক সময়: ৬০ মিনিট

মনোযোগ কীভাবে কাজ করে

স্কেলড ডট-প্রোডাক্ট অ্যাটেনশন তিনটি ইনপুট— কোয়েরি , কী এবং ভ্যালু —গ্রহণ করে এবং গণনা করে:

Attention(Q, K, V) = softmax(Q K^T / sqrt(d_k)) V

স্কেলিং ফ্যাক্টর 1 / sqrt(d_k) হেড ডাইমেনশন বাড়ার সাথে সাথে ডট প্রোডাক্টগুলোকে অতিরিক্ত বড় হওয়া থেকে বিরত রাখে, যা সফটম্যাক্সকে এমন অঞ্চলে ঠেলে দেবে যেখানে এর গ্রেডিয়েন্টগুলো খুবই ক্ষুদ্র।

চারটি ধাপ, যার প্রতিটি পরবর্তী ধাপকে চালিত করে:

  1. স্কোরQ @ KT : ডট প্রোডাক্ট পরিমাপ করে যে প্রতিটি কোয়েরি পজিশন প্রতিটি কী পজিশনের প্রতি কতটা মনোযোগ দেবে।
  2. স্কেল + সফটম্যাক্সsoftmax(scores / sqrt(d)) : স্কেলিং গ্রেডিয়েন্ট অদৃশ্য হওয়া রোধ করে, এবং সফটম্যাক্স স্কোরগুলিকে অ্যাটেনশন ওয়েটে রূপান্তরিত করে যেগুলির যোগফল ১ হয়।
  3. Attendweights @ V : ভ্যালু ভেক্টরগুলোর ওয়েটেড সাম প্রতিটি কোয়েরি পজিশনের জন্য আউটপুট তৈরি করে।
  4. আউটপুট — Q-এর অনুরূপ আকৃতি: প্রতিটি কোয়েরি পজিশন এখন তার অ্যাটেন্ড করা পজিশনগুলো থেকে তথ্য বহন করে।

২. শুরু করার আগে

আপনার প্রকল্প নির্বাচন করুন

গুগল ক্লাউড কনসোলে , বিলিং সক্ষম করা আছে এমন একটি প্রজেক্ট নির্বাচন করুন বা তৈরি করুন।

ওপেন ক্লাউড শেল

একটি ক্লাউড শেল সেশন শুরু করতে অ্যাক্টিভেট ক্লাউড শেল (কনসোলের উপরের ডানদিকে থাকা টার্মিনাল আইকন)-এ ক্লিক করুন, তারপর এটিকে আপনার প্রজেক্টে নির্দেশ করুন:

gcloud config set project <YOUR_PROJECT_ID>

জিপিইউ পরিবেশ প্রস্তুত করুন

ক্লাউড শেলে নিম্নলিখিতটি চালান। কোডল্যাব ১-এ জিপিইউ কোটা এবং জোনের প্রয়োজনীয়তা সহ প্রতিটি কমান্ড বিস্তারিতভাবে ব্যাখ্যা করা হয়েছে।

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:// :8884 খুলুন http:// :8884 , টোকেনটি পেস্ট করুন এবং /workspace এ একটি নতুন পাইথন ৩ নোটবুক তৈরি করুন। এই কোডল্যাবের প্রতিটি কোড ব্লক সেই নোটবুকের একটি সেলে যাবে।

এই কোডল্যাবের জন্য যা যা প্রয়োজন তা ইনস্টল করুন।

!pip install --quiet matplotlib flax

জিপিইউ সেট আপ এবং যাচাই করুন

JAX ইম্পোর্ট করুন এবং নিশ্চিত করুন যে ডিফল্ট ব্যাকএন্ডটি একটি GPU। এই সেলটি block_tree , show_table , এবং show_bars কেও সংজ্ঞায়িত করে, যেগুলো পরবর্তী প্রতিটি ধাপে ডিভাইসের কাজের জন্য অপেক্ষা করতে এবং ফলাফল রেন্ডার করতে ব্যবহৃত হয়।

import os
os.environ["LD_LIBRARY_PATH"] = "/usr/local/nvidia/lib64:" + os.environ.get("LD_LIBRARY_PATH", "")

import html
import math
import time
from functools import partial

from IPython.display import HTML, display
import matplotlib.pyplot as plt
import numpy as np

import jax
import jax.numpy as jnp


devices = jax.devices()
gpu_devices = [d for d in devices if d.platform == "gpu"]
device = gpu_devices[0] if gpu_devices else None

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"Using GPU:       {device}")


def block_tree(tree):
    """Wait until a PyTree of JAX arrays is ready on device."""
    return jax.block_until_ready(tree)


def show_table(headers, rows, title=None, aligns=None):
    """Render rows as an HTML table."""
    aligns = aligns or ["left"] * len(headers)
    parts = ["<div style='font-family: system-ui; max-width: 980px;'>"]
    if title:
        parts.append(f"<h4 style='margin: 0 0 8px 0;'>{html.escape(title)}</h4>")
    parts.append("<table style='border-collapse: collapse; width: 100%; font-size: 13px;'>")
    parts.append("<thead><tr>")
    for h, a in zip(headers, aligns):
        parts.append(
            f"<th style='text-align:{a}; border-bottom:1px solid #d0d7de; padding:6px;'>"
            f"{html.escape(str(h))}</th>"
        )
    parts.append("</tr></thead><tbody>")
    for row in rows:
        parts.append("<tr>")
        for cell, a in zip(row, aligns):
            parts.append(
                f"<td style='text-align:{a}; border-bottom:1px solid #eef1f4; padding:6px;'>"
                f"{html.escape(str(cell))}</td>"
            )
        parts.append("</tr>")
    parts.append("</tbody></table></div>")
    display(HTML("".join(parts)))


def show_bars(rows, title, unit="", lower_is_better=False):
    """Render (label, value) pairs as a horizontal bar chart in HTML."""
    max_value = max(float(value) for _, value in rows) or 1.0
    color = "#1a7f37" if not lower_is_better else "#0969da"
    parts = ["<div style='font-family: Arial, sans-serif; max-width: 760px;'>"]
    parts.append(f"<h4 style='margin: 0 0 8px 0;'>{html.escape(title)}</h4>")
    for label, value in rows:
        width = max(3, 100 * float(value) / max_value)
        parts.append(
            "<div style='display:grid; grid-template-columns: 190px 1fr 130px; gap: 8px; "
            "align-items:center; margin: 6px 0;'>"
            f"<div style='font-size:13px;'>{html.escape(str(label))}</div>"
            "<div style='background:#f6f8fa; border-radius:6px; overflow:hidden; height:22px;'>"
            f"<div style='height:22px; width:{width:.1f}%; background:{color};'></div></div>"
            f"<div style='font-size:13px; font-variant-numeric: tabular-nums;'>{float(value):,.1f} {html.escape(unit)}</div>"
            "</div>"
        )
    parts.append(
        f"<div style='font-size:12px; color:#57606a;'>"
        f"{'Lower' if lower_is_better else 'Higher'} is better.</div></div>"
    )
    display(HTML("".join(parts)))

আপনি একটি JAX ভার্সন, ডিফল্ট ব্যাকএন্ড হিসেবে gpu , CUDA ডিভাইসগুলোর একটি তালিকা এবং কোডল্যাবের বাকি অংশে ব্যবহৃত GPU-টি দেখতে পাবেন।

৩. Q, K, এবং V টেস্ট অ্যারে তৈরি করুন

ফিউজড অ্যাটেনশন কার্নেলগুলো মেমরিতে ইনপুট অ্যারেগুলো কীভাবে সাজানো আছে সে বিষয়ে যত্নশীল, তাই যেকোনো অ্যাটেনশন কোড লেখার আগে JAX-এর প্রত্যাশিত লেআউটে Q, K, এবং V তৈরি করুন। jax.nn.dot_product_attention নিম্নলিখিত লেআউটটি প্রত্যাশা করে:

ম্লান

অর্থ

আমাদের ডিফল্ট

বি

ব্যাচের আকার

টি

কোয়েরি ক্রমের দৈর্ঘ্য

১২৮

এস

কী/মান ক্রমের দৈর্ঘ্য

১২৮ (আত্ম-মনোযোগের জন্য T-এর সমান)

এন

অ্যাটেনশন হেডের সংখ্যা

এইচ

মাথা প্রতি মাত্রা

৬৪

সরল বাস্তবায়নের জন্য, আপনি র‍্যান্ডম মান দিয়ে ইনিশিয়ালাইজ করে float32 থেকে শুরু করেন।

BATCH = 4
SEQ_LEN = 128
NUM_HEADS = 8
HEAD_DIM = 64

key = jax.random.key(0)
k1, k2, k3 = jax.random.split(key, 3)

q = jax.random.normal(k1, (BATCH, SEQ_LEN, NUM_HEADS, HEAD_DIM), dtype=jnp.float32)
k = jax.random.normal(k2, (BATCH, SEQ_LEN, NUM_HEADS, HEAD_DIM), dtype=jnp.float32)
v = jax.random.normal(k3, (BATCH, SEQ_LEN, NUM_HEADS, HEAD_DIM), dtype=jnp.float32)

q, k, v = jax.device_put((q, k, v), device)

show_table(
    ["Array", "Shape", "Dtype", "Layout"],
    [
        ("Q (query)", q.shape, q.dtype, "(B, T, N, H)"),
        ("K (key)", k.shape, k.dtype, "(B, S, N, H)"),
        ("V (value)", v.shape, v.dtype, "(B, S, N, H)"),
    ],
    title="Attention inputs on GPU",
)

আপনি একটি টেবিল দেখতে পাবেন যেখানে প্রতিটি অ্যারের জন্য একটি করে সারি থাকবে, এবং প্রতিটি সারিতে (4, 128, 8, 64) আকৃতি ও float32 ডেটাটাইপ রিপোর্ট করা থাকবে। jax.device_put কলটি তিনটি অ্যারেকেই আপনার সেটআপ সেলে নির্বাচিত GPU-তে পিন করে দেয়, তাই পরবর্তী বেঞ্চমার্কগুলিতে কোনো হোস্ট ট্রান্সফার পরিমাপ করা হয় না।

৪. গোড়া থেকে মনোযোগ বাস্তবায়ন করুন

ফিউজড কার্নেল ব্যবহার করার আগে, সূত্রটি সরাসরি লিখে ফেলুন যাতে আপনি দেখতে পারেন কার্নেলটি ঠিক কী প্রতিস্থাপন করছে। এই ইমপ্লিমেন্টেশনটি ক্রমানুসারে চারটি ধাপ অনুসরণ করে: স্কোর গণনা, স্কেল, সফটম্যাক্স, এবং মান দ্বারা গুণ। আপনি হেড এবং সিকোয়েন্স ডাইমেনশন ট্রান্সপোজ করেন, ফলে ম্যাটমালটি সিকোয়েন্স অ্যাক্সিসের উপর কাজ করে।

এটি সঠিকভাবে কাজ করে, কিন্তু এটি স্কোর ম্যাট্রিক্সমাল্টিপল, সফটম্যাক্স, ভ্যালু ম্যাট্রিক্সমাল্টিপল-এর ​​জন্য তিনটি পৃথক জিপিইউ কার্নেল লঞ্চ করে এবং সম্পূর্ণ (B, N, T, S) অ্যাটেনশন ওয়েট ম্যাট্রিক্সটিকে জিপিইউ মেমরিতে মেটেরিয়ালাইজ করে।

def naive_attention(q, k, v):
    """Scaled dot-product attention from scratch."""
    scale = 1.0 / math.sqrt(q.shape[-1])

    # (B, T, N, H) to (B, N, T, H) so matmul runs over the T/S axis
    q_t = jnp.transpose(q, (0, 2, 1, 3))
    k_t = jnp.transpose(k, (0, 2, 1, 3))
    v_t = jnp.transpose(v, (0, 2, 1, 3))

    # Score: (B, N, T, H) @ (B, N, H, S) to (B, N, T, S)
    scores = jnp.matmul(q_t, jnp.transpose(k_t, (0, 1, 3, 2))) * scale
    weights = jax.nn.softmax(scores, axis=-1)

    # Attend: (B, N, T, S) @ (B, N, S, H) to (B, N, T, H)
    out_t = jnp.matmul(weights, v_t)

    # Back to (B, T, N, H)
    return jnp.transpose(out_t, (0, 2, 1, 3))


naive_out = block_tree(naive_attention(q, k, v))

show_table(
    ["", "Value"],
    [
        ("Output shape", str(naive_out.shape)),
        ("Output dtype", str(naive_out.dtype)),
    ],
    title="Naive attention",
)

আপনি আউটপুটের আকৃতি (4, 128, 8, 64) এবং ডেটা টাইপ float32 দেখতে পাবেন, যার আকৃতি Q-এর সমান, এবং ফর্মুলার আউটপুট ধাপে ঠিক এটাই প্রতিশ্রুতি দেওয়া হয়েছিল।

৫. ডট_প্রোডাক্ট_অ্যাটেনশন-এ পরিবর্তন করুন

jax.nn.dot_product_attention , score, scale, softmax, এবং attend ধাপগুলোকে একটি একক অপারেশনে একীভূত করে। এর ফলে JAX এবং XLA মেমরি অ্যাক্সেস প্যাটার্নকে অপ্টিমাইজ করতে পারে। বিশেষত, সিকোয়েন্স দীর্ঘ হলে তারা সম্পূর্ণ অ্যাটেনশন ওয়েট ম্যাট্রিক্সকে বাস্তবায়িত করা এড়াতে পারে।

ডিফল্ট implementation=None করা থাকলে, JAX স্বয়ংক্রিয়ভাবে সেরা উপলব্ধ ব্যাকএন্ডটি বেছে নেয়। cuDNN উপলব্ধ এবং সামঞ্জস্যপূর্ণ ইনপুটযুক্ত কোনো GPU-তে এটি আগে থেকেই cuDNN ব্যবহার করতে পারে। অন্যান্য হার্ডওয়্যারে এটি XLA-তে ফিরে যায়।

sdpa_out = block_tree(jax.nn.dot_product_attention(q, k, v))

max_diff = float(jnp.max(jnp.abs(naive_out - sdpa_out)))

show_table(
    ["", "Value"],
    [
        ("Output shape", str(sdpa_out.shape)),
        ("Output dtype", str(sdpa_out.dtype)),
        ("Max |naive − SDPA|", f"{max_diff:.2e}"),
        ("Outputs close (atol=1e-3)", str(bool(jnp.allclose(naive_out, sdpa_out, atol=1e-3)))),
    ],
    title="JAX SDPA vs naive",
)

আপনি সরল সংস্করণের মতোই একই আকৃতি ও ডেটাটাইপ, একটি ক্ষুদ্র সর্বোচ্চ পরম পার্থক্য এবং atol=1e-3 নৈকট্য যাচাইয়ের জন্য True দেখতে পাবেন।

৬. cuDNN ফিউজড ব্যাকএন্ডকে জোরপূর্বক প্রয়োগ করুন

JAX-কে ব্যাকএন্ড বেছে নিতে দেওয়া সুবিধাজনক, কিন্তু বাইরে থেকে বোঝা যায় না আসলে কোন কার্নেলটি চলেছে। implementation="cudnn" সেট করলে NVIDIA-র cuDNN ফিউজড অ্যাটেনশন কার্নেলগুলো ব্যবহৃত হয়। এগুলো হলো হাতে অপ্টিমাইজ করা GPU কার্নেল, যা অপ্টিমাইজ করা মেমরি অ্যাক্সেস প্যাটার্নসহ সম্পূর্ণ অ্যাটেনশন গণনাকে একটিমাত্র কার্নেল লঞ্চের মধ্যে একীভূত করে।

লক্ষ্য করুন যে cuDNN ফিউজড অ্যাটেনশনের কিছু হার্ডওয়্যার এবং শেপ রিকোয়ারমেন্ট রয়েছে, যেমন GPU কম্পিউট ক্যাপাবিলিটি যা অবশ্যই >= 8.0 (অ্যাম্পিয়ার বা নতুন) হতে হবে, অথবা ইনপুট ডেটাটাইপ হিসেবে float16 বা bfloat16 হতে হবে। যদি এই রিকোয়ারমেন্টগুলো পূরণ না হয় এবং আপনি implementation="cudnn" সেট করেন, তাহলে JAX নীরবে ফলব্যাক না করে একটি এরর দেখায়। এই কারণেই নিচের কোডটি bfloat16 এ কাস্ট করে এবং কলটিকে try / except মধ্যে রাখে: এটি HAS_CUDNN_SDPA সেট করে, যাতে পরবর্তী প্রতিটি ধাপ জানতে পারে যে এই মেশিনে cuDNN পাথটি উপলব্ধ আছে কিনা।

q_bf16 = q.astype(jnp.bfloat16)
k_bf16 = k.astype(jnp.bfloat16)
v_bf16 = v.astype(jnp.bfloat16)

HAS_CUDNN_SDPA = False

try:
    cudnn_out = block_tree(
        jax.nn.dot_product_attention(q_bf16, k_bf16, v_bf16, implementation="cudnn")
    )
    HAS_CUDNN_SDPA = True

    xla_bf16_out = block_tree(
        jax.nn.dot_product_attention(q_bf16, k_bf16, v_bf16, implementation="xla")
    )
    max_diff = float(jnp.max(jnp.abs(
        cudnn_out.astype(jnp.float32) - xla_bf16_out.astype(jnp.float32)
    )))

    show_table(
        ["", "Value"],
        [
            ("Output shape", str(cudnn_out.shape)),
            ("Output dtype", str(cudnn_out.dtype)),
            ("Max |cuDNN − XLA| (both bf16)", f"{max_diff:.2e}"),
            ("Outputs close (rtol=1e-2, atol=1e-2)", str(bool(jnp.allclose(cudnn_out, xla_bf16_out, rtol=1e-2, atol=1e-2)))),
        ],
        title="cuDNN fused attention",
    )

except Exception as e:
    print(f"cuDNN SDPA not available on this GPU: {e}")
    print("Continuing with XLA backend only.")

যদি আপনার GPU, cuDNN SDPA চালাতে না পারে, তাহলে কোডটি তার কারণ প্রিন্ট করে এবং কোডল্যাবটি XLA ব্যাকএন্ডে চলতে থাকে।

৭. কার্যকারণ মাস্কিং যোগ করুন

অটোরেগ্রেসিভ মডেলে (GPT-স্টাইল ডিকোডার), প্রতিটি পজিশন কেবল তার পূর্ববর্তী পজিশনগুলোকেই অ্যাটেন্ড করতে পারে। is_causal=True সেট করলে এই লোয়ার-ট্রায়াঙ্গুলার মাস্কটি ফিউজড কার্নেলের ভিতরে প্রয়োগ হয়, ফলে আপনাকে নিজে থেকে মাস্ক ম্যাট্রিক্স তৈরি করতে হয় না।

নিচের কোডটি মাস্ক সহ এবং মাস্ক ছাড়া দুইবার অ্যাটেনশন রান করে এবং মাস্কটি কী পরিবর্তন করেছে তা দেখানোর জন্য দুটি পজিশন তুলনা করে।

causal_out = block_tree(
    jax.nn.dot_product_attention(q, k, v, is_causal=True)
)

# With causal masking, the last position attends to all positions.
# The first position attends only to itself.
nocausal_out = block_tree(
    jax.nn.dot_product_attention(q, k, v, is_causal=False)
)

# First position should differ
first_pos_diff = float(jnp.max(jnp.abs(causal_out[:, 0] - nocausal_out[:, 0])))
# Last position should be the same
last_pos_diff = float(jnp.max(jnp.abs(causal_out[:, -1] - nocausal_out[:, -1])))

show_table(
    ["Position", "Max diff (causal vs full)", "Expected"],
    [
        ("First (t=0)", f"{first_pos_diff:.4f}", "Large — causal restricts to self only"),
        ("Last (t=T-1)", f"{last_pos_diff:.2e}", "~0 — attends to all positions either way"),
    ],
    title="Causal masking effect on attention output",
)

আপনি প্রথম অবস্থানে একটি বড় পার্থক্য এবং শেষ অবস্থানে প্রায়-শূন্য পার্থক্য দেখতে পাবেন। এই অপ্রতিসাম্যের কারণ হলো, ০ নম্বর অবস্থানের মাস্কটি তার নিজের কী ছাড়া বাকি সব কী হারিয়ে ফেলে, অথচ শেষ অবস্থানটি ইতিমধ্যেই সম্পূর্ণ সিকোয়েন্সটি দেখতে পায়।

৮. কিছু বেঞ্চমার্ক চালানো

এখন আপনার কাছে একই ফাংশন গণনা করার তিনটি উপায় আছে। কোনটি বেছে নেবেন তা সম্পূর্ণরূপে সমস্যার ধরনের উপর নির্ভর করে এবং আপনি এর মান নির্ণয় করতে পারেন।

ডিফল্ট আকারে ভ্যারিয়েন্টগুলোর সময় পরিমাপ করুন।

def benchmark_attention(fn, q, k, v, warmup=3, repeats=50):
    """Time an attention function. Returns median milliseconds per call."""
    jit_fn = jax.jit(fn)

    for _ in range(warmup):
        block_tree(jit_fn(q, k, v))

    times = []
    for _ in range(repeats):
        start = time.perf_counter()
        block_tree(jit_fn(q, k, v))
        times.append((time.perf_counter() - start) * 1000)

    return np.median(times)


t_naive = benchmark_attention(naive_attention, q, k, v)
t_sdpa = benchmark_attention(
    lambda q, k, v: jax.nn.dot_product_attention(q, k, v, implementation="xla"),
    q, k, v,
)

results = [
    ("Naive (matmul + softmax + matmul)", f"{t_naive:.2f}"),
    ("SDPA (XLA, float32)", f"{t_sdpa:.2f}"),
]
bar_data = [
    ("Naive", t_naive),
    ("SDPA XLA f32", t_sdpa),
]

if HAS_CUDNN_SDPA:
    t_sdpa_bf16 = benchmark_attention(
        lambda q, k, v: jax.nn.dot_product_attention(q, k, v, implementation="xla"),
        q_bf16, k_bf16, v_bf16,
    )
    t_cudnn = benchmark_attention(
        lambda q, k, v: jax.nn.dot_product_attention(q, k, v, implementation="cudnn"),
        q_bf16, k_bf16, v_bf16,
    )
    results.append(("SDPA (XLA, bfloat16)", f"{t_sdpa_bf16:.2f}"))
    results.append(("SDPA (cuDNN, bfloat16)", f"{t_cudnn:.2f}"))
    bar_data.append(("SDPA XLA bf16", t_sdpa_bf16))
    bar_data.append(("SDPA cuDNN bf16", t_cudnn))

show_table(
    ["Implementation", "Median ms/call"],
    results,
    title=f"Attention timing — B={BATCH}, T={SEQ_LEN}, N={NUM_HEADS}, H={HEAD_DIM}",
    aligns=["left", "right"],
)
show_bars(bar_data, "Attention latency (ms per call)", "ms", lower_is_better=True)

এই ছোট আকারের সমস্যার জন্য, সমস্ত ইমপ্লিমেন্টেশনই খুব দ্রুত এবং বেশিরভাগ ক্ষেত্রেই ওভারহেড-সীমিত, তাই সরল সংস্করণটি প্রতিযোগিতামূলক। cuDNN bf16 সামান্য দ্রুততম হওয়া উচিত, কিন্তু দীর্ঘতর সিকোয়েন্স দৈর্ঘ্যে ফিউজড অ্যাটেনশনের সুবিধা সাধারণত আরও স্পষ্ট হয়ে ওঠে, যেখানে সম্পূর্ণ অ্যাটেনশন ম্যাট্রিক্স এড়িয়ে চলাটা গুরুত্বপূর্ণ।

ক্রমের দৈর্ঘ্য সুইপ করুন

সিকোয়েন্সের দৈর্ঘ্য বাড়ার সাথে সাথে ফিউজড অ্যাটেনশন কার্নেলের সুবিধাও বৃদ্ধি পায়। এর সাধারণ ইমপ্লিমেন্টেশন GPU মেমরিতে একটি (B, N, T, S) অ্যাটেনশন ম্যাট্রিক্স তৈরি করে, যার জন্য O(T²) মেমরি প্রয়োজন হয়। cuDNN FlashAttention-এর মতো ফিউজড কার্নেলগুলো কম্পিউটেশনকে এমনভাবে ভাগ করে যে তারা কখনোই সম্পূর্ণ ম্যাট্রিক্স তৈরি করে না, ফলে মেমরি O(T) থাকে।

নিচের সুইপটি ৬৪ থেকে ১০২৪ পর্যন্ত সিকোয়েন্স দৈর্ঘ্যের মধ্যে প্রতিটি ইমপ্লিমেন্টেশনের সময় পরিমাপ করে।

SEQ_LENS = [64, 128, 256, 512, 1024]

sweep_results = []
for sl in SEQ_LENS:
    rk = jax.random.key(sl)
    rk1, rk2, rk3 = jax.random.split(rk, 3)

    q_s = jax.random.normal(rk1, (BATCH, sl, NUM_HEADS, HEAD_DIM), dtype=jnp.float32)
    k_s = jax.random.normal(rk2, (BATCH, sl, NUM_HEADS, HEAD_DIM), dtype=jnp.float32)
    v_s = jax.random.normal(rk3, (BATCH, sl, NUM_HEADS, HEAD_DIM), dtype=jnp.float32)
    q_s, k_s, v_s = jax.device_put((q_s, k_s, v_s), device)

    q_sb = q_s.astype(jnp.bfloat16)
    k_sb = k_s.astype(jnp.bfloat16)
    v_sb = v_s.astype(jnp.bfloat16)

    row = {"seq_len": sl}

    row["naive_ms"] = benchmark_attention(
        naive_attention,
        q_s, k_s, v_s,
        warmup=2,
        repeats=20,
    )

    row["sdpa_xla_f32_ms"] = benchmark_attention(
        lambda q, k, v: jax.nn.dot_product_attention(q, k, v, implementation="xla"),
        q_s, k_s, v_s,
        warmup=2,
        repeats=20,
    )

    row["sdpa_xla_bf16_ms"] = benchmark_attention(
        lambda q, k, v: jax.nn.dot_product_attention(q, k, v, implementation="xla"),
        q_sb, k_sb, v_sb,
        warmup=2,
        repeats=20,
    )

    if HAS_CUDNN_SDPA:
        row["cudnn_bf16_ms"] = benchmark_attention(
            lambda q, k, v: jax.nn.dot_product_attention(q, k, v, implementation="cudnn"),
            q_sb, k_sb, v_sb,
            warmup=2,
            repeats=20,
        )

    sweep_results.append(row)


headers = ["Seq len", "Naive (ms)", "SDPA XLA f32 (ms)", "SDPA XLA bf16 (ms)"]
if HAS_CUDNN_SDPA:
    headers.append("cuDNN bf16 (ms)")

table_rows = []
for r in sweep_results:
    row = [
        r["seq_len"],
        f"{r['naive_ms']:.2f}",
        f"{r['sdpa_xla_f32_ms']:.2f}",
        f"{r['sdpa_xla_bf16_ms']:.2f}",
    ]

    if HAS_CUDNN_SDPA:
        row.append(f"{r['cudnn_bf16_ms']:.2f}")

    table_rows.append(row)

show_table(
    headers,
    table_rows,
    title=f"Sequence-length sweep — B={BATCH}, N={NUM_HEADS}, H={HEAD_DIM}",
    aligns=["right"] * len(headers),
)

এই প্রক্রিয়াটিতে কিছুটা সময় লাগে, কারণ প্রতিটি সিকোয়েন্স দৈর্ঘ্যের সময় পরিমাপ করার আগে এটি তিন বা চারটি আলাদা ইমপ্লিমেন্টেশন কম্পাইল করে। সবশেষে, ৬৪ থেকে ১০২৪ পর্যন্ত প্রতিটি সিকোয়েন্স দৈর্ঘ্যের জন্য আপনার একটি করে টেবিল সারি থাকা উচিত।

এই প্লটটি নেটিভ ইমপ্লিমেন্টেশন, float32 ও bfloat16 ফরম্যাটে JAX/XLA SDPA, এবং bfloat16 ফরম্যাটে cuDNN ফিউজড অ্যাটেনশনের ক্ষেত্রে সিকোয়েন্সের দৈর্ঘ্যের সাথে অ্যাটেনশন ল্যাটেন্সির পরিবর্তনের তুলনা করে।

fig, ax = plt.subplots(figsize=(8, 5))
seq_lens = [r["seq_len"] for r in sweep_results]

ax.plot(
    seq_lens,
    [r["naive_ms"] for r in sweep_results],
    "o-",
    label="Naive",
    color="#d1242f",
)

ax.plot(
    seq_lens,
    [r["sdpa_xla_f32_ms"] for r in sweep_results],
    "s-",
    label="SDPA XLA f32",
    color="#0969da",
)

ax.plot(
    seq_lens,
    [r["sdpa_xla_bf16_ms"] for r in sweep_results],
    "d-",
    label="SDPA XLA bf16",
    color="#8250df",
)

if HAS_CUDNN_SDPA:
    ax.plot(
        seq_lens,
        [r["cudnn_bf16_ms"] for r in sweep_results],
        "^-",
        label="cuDNN bf16",
        color="#1a7f37",
    )

ax.set_xlabel("Sequence length")
ax.set_ylabel("Median ms per call")
ax.set_title("Attention latency vs sequence length")
ax.legend()
ax.grid(True, alpha=0.25)
ax.set_xticks(seq_lens)

fig.tight_layout()
plt.show()

প্রতিটি বাস্তবায়নের জন্য একটি করে লাইন দেখা উচিত। এবং অনুক্রমটি দীর্ঘ হওয়ার সাথে সাথে যে ব্যবধানটি বাড়তে থাকে, সেটিও আপনার দেখা উচিত।

ব্যাচ সাইজ সুইপ করুন

বৃহত্তর ব্যাচগুলো কার্নেল লঞ্চের ওভারহেড পুষিয়ে দেয় এবং জিপিইউ-এর ব্যবহার উন্নত করে, যতক্ষণ না জিপিইউ মেমরি একটি বাধা হয়ে দাঁড়ায়। নিচের সুইপটিতে সিকোয়েন্সের দৈর্ঘ্য ২৫৬-এ স্থির রেখে ব্যাচের আকার পরিবর্তন করা হয়েছে।

BATCH_SIZES = [1, 2, 4, 8, 16]
SWEEP_SEQ = 256

batch_results = []
for bs in BATCH_SIZES:
    rk = jax.random.key(bs + 100)
    rk1, rk2, rk3 = jax.random.split(rk, 3)

    q_b = jax.random.normal(rk1, (bs, SWEEP_SEQ, NUM_HEADS, HEAD_DIM), dtype=jnp.float32)
    k_b = jax.random.normal(rk2, (bs, SWEEP_SEQ, NUM_HEADS, HEAD_DIM), dtype=jnp.float32)
    v_b = jax.random.normal(rk3, (bs, SWEEP_SEQ, NUM_HEADS, HEAD_DIM), dtype=jnp.float32)
    q_b, k_b, v_b = jax.device_put((q_b, k_b, v_b), device)

    q_bb = q_b.astype(jnp.bfloat16)
    k_bb = k_b.astype(jnp.bfloat16)
    v_bb = v_b.astype(jnp.bfloat16)

    row = {"batch": bs}

    row["sdpa_xla_f32_ms"] = benchmark_attention(
        lambda q, k, v: jax.nn.dot_product_attention(q, k, v, implementation="xla"),
        q_b, k_b, v_b,
        warmup=2,
        repeats=20,
    )

    row["sdpa_xla_bf16_ms"] = benchmark_attention(
        lambda q, k, v: jax.nn.dot_product_attention(q, k, v, implementation="xla"),
        q_bb, k_bb, v_bb,
        warmup=2,
        repeats=20,
    )

    if HAS_CUDNN_SDPA:
        row["cudnn_bf16_ms"] = benchmark_attention(
            lambda q, k, v: jax.nn.dot_product_attention(q, k, v, implementation="cudnn"),
            q_bb, k_bb, v_bb,
            warmup=2,
            repeats=20,
        )

    batch_results.append(row)


headers = ["Batch size", "SDPA XLA f32 (ms)", "SDPA XLA bf16 (ms)"]
if HAS_CUDNN_SDPA:
    headers.append("cuDNN bf16 (ms)")

table_rows = []
for r in batch_results:
    row = [
        r["batch"],
        f"{r['sdpa_xla_f32_ms']:.2f}",
        f"{r['sdpa_xla_bf16_ms']:.2f}",
    ]

    if HAS_CUDNN_SDPA:
        row.append(f"{r['cudnn_bf16_ms']:.2f}")

    table_rows.append(row)

show_table(
    headers,
    table_rows,
    title=f"Batch-size sweep — T={SWEEP_SEQ}, N={NUM_HEADS}, H={HEAD_DIM}",
    aligns=["right"] * len(headers),
)


fig, ax = plt.subplots(figsize=(8, 5))
batches = [r["batch"] for r in batch_results]

ax.plot(
    batches,
    [r["sdpa_xla_f32_ms"] for r in batch_results],
    "s-",
    label="SDPA XLA f32",
    color="#0969da",
)

ax.plot(
    batches,
    [r["sdpa_xla_bf16_ms"] for r in batch_results],
    "d-",
    label="SDPA XLA bf16",
    color="#8250df",
)

if HAS_CUDNN_SDPA:
    ax.plot(
        batches,
        [r["cudnn_bf16_ms"] for r in batch_results],
        "^-",
        label="cuDNN bf16",
        color="#1a7f37",
    )

ax.set_xlabel("Batch size")
ax.set_ylabel("Median ms per call")
ax.set_title("Attention latency vs batch size")
ax.legend()
ax.grid(True, alpha=0.25)
ax.set_xticks(batches)

fig.tight_layout()
plt.show()

আপনি ১ থেকে ১৬ পর্যন্ত প্রতিটি ব্যাচ সাইজের জন্য একটি করে এন্ট্রি সহ একটি টেবিল এবং একটি প্লট পাবেন। এই পর্যালোচনার বাইরে রাখা হয়েছে সরল ইমপ্লিমেন্টেশনটি, তাই উভয় আউটপুটই কেবল তিনটি ফিউজড পাথের তুলনা করে।

৯. MHA, GQA এবং MQA-এর মধ্যে তুলনা করুন।

এখন পর্যন্ত প্রতিটি হেড তার নিজস্ব কী এবং ভ্যালু বহন করে আসছে। কিন্তু ডিকোডারগুলো প্রায়শই সেই প্রতিসাম্য ভেঙে দেয়, কারণ ইনফারেন্সের সময় কী-ভ্যালু ক্যাশই মেমোরিতে প্রাধান্য বিস্তার করে।

মাল্টি-হেড অ্যাটেনশন (MHA) প্রতিটি হেডকে তার নিজস্ব Q, K, এবং V প্রজেকশন প্রদান করে। গ্রুপড-কোয়েরি অ্যাটেনশন (GQA) এবং মাল্টি-কোয়েরি অ্যাটেনশন (MQA) ইনফারেন্সের সময় মেমরি এবং গণনা বাঁচাতে KV হেডের সংখ্যা হ্রাস করে।

jax.nn.dot_product_attention এই তিনটিই পরিচালনা করে এবং যখন K < N হয়, তখন KV হেডগুলো স্বয়ংক্রিয়ভাবে ব্রডকাস্ট করা হয়।

rk = jax.random.key(42)
rk1, rk2, rk3, rk4, rk5 = jax.random.split(rk, 5)

q_mha = jax.random.normal(rk1, (2, 64, 8, 64), dtype=jnp.float32)

# MHA: 8 KV heads
k_mha = jax.random.normal(rk2, (2, 64, 8, 64), dtype=jnp.float32)
v_mha = jax.random.normal(rk3, (2, 64, 8, 64), dtype=jnp.float32)

# GQA: 2 KV heads (each shared by 4 query heads)
k_gqa = jax.random.normal(rk2, (2, 64, 2, 64), dtype=jnp.float32)
v_gqa = jax.random.normal(rk3, (2, 64, 2, 64), dtype=jnp.float32)

# MQA: 1 KV head (shared by all 8 query heads)
k_mqa = jax.random.normal(rk4, (2, 64, 1, 64), dtype=jnp.float32)
v_mqa = jax.random.normal(rk5, (2, 64, 1, 64), dtype=jnp.float32)

out_mha = block_tree(jax.nn.dot_product_attention(q_mha, k_mha, v_mha))
out_gqa = block_tree(jax.nn.dot_product_attention(q_mha, k_gqa, v_gqa))
out_mqa = block_tree(jax.nn.dot_product_attention(q_mha, k_mqa, v_mqa))

show_table(
    ["Pattern", "Q shape", "K shape", "V shape", "Output shape"],
    [
        ("MHA", q_mha.shape, k_mha.shape, v_mha.shape, out_mha.shape),
        ("GQA", q_mha.shape, k_gqa.shape, v_gqa.shape, out_gqa.shape),
        ("MQA", q_mha.shape, k_mqa.shape, v_mqa.shape, out_mqa.shape),
    ],
    title="Multi-head attention variants — all should produce the same output shape",
)

যদিও K এবং V শেপ ৮টি হেড থেকে কমে ২টি ও পরে ১টি হয়, তবুও তিনটি সারিরই আউটপুট শেপ একই (2, 64, 8, 64) হওয়া উচিত। মূল বিষয়টি হলো: আপনি অ্যাটেনশনের পরবর্তী ধাপে কোনো কিছু পরিবর্তন না করেই KV ক্যাশে বাদ দিতে পারেন।

১০. এনভিডিয়া ট্রান্সফরমারইঞ্জিন এবং এফপি৮ এর বেঞ্চমার্ক

এনভিডিয়ার TransformerEngine , এনভিডিয়া জিপিইউ-এর জন্য অপ্টিমাইজ করা ফিউজড অ্যাটেনশন মডিউল সরবরাহ করে। JAX ইন্টিগ্রেশনটি Flax Linen-স্টাইলের মডিউল ব্যবহার করে (NNX নয়), তাই মডিউলটি একবার ইনিশিয়ালাইজ করা হয় এবং তারপর এর ভ্যারিয়েবলগুলোসহ প্রয়োগ করা হয়।

TransformerEngine-এর সাথে সুইপ সিকোয়েন্সের দৈর্ঘ্য

এই সেলটি প্রথমে NVIDIA TransformerEngine উপলব্ধ আছে কিনা তা যাচাই করে, এবং তারপর একই ওয়ার্কলোডে XLA ও cuDNN ব্যাকএন্ডসহ JAX SDPA-এর বিপরীতে বিভিন্ন সিকোয়েন্স দৈর্ঘ্যে এর bf16 কজাল DotProductAttention বেঞ্চমার্ক পরীক্ষা করে।

TE_SEQ_LENS = [128, 256, 512, 1024, 2048]
TE_BATCH = BATCH

HAS_TE = False

try:
    import transformer_engine.jax as te
    import transformer_engine.jax.flax as te_flax
    HAS_TE = True
except ImportError:
    print("TransformerEngine not installed — skipping TE sections.")

if HAS_TE:
    te_results = []

    for sl in TE_SEQ_LENS:
        rk = jax.random.key(sl + 1000)
        rk1, rk2, rk3 = jax.random.split(rk, 3)

        q_te = jax.random.normal(
            rk1, (TE_BATCH, sl, NUM_HEADS, HEAD_DIM), dtype=jnp.bfloat16
        )
        k_te = jax.random.normal(
            rk2, (TE_BATCH, sl, NUM_HEADS, HEAD_DIM), dtype=jnp.bfloat16
        )
        v_te = jax.random.normal(
            rk3, (TE_BATCH, sl, NUM_HEADS, HEAD_DIM), dtype=jnp.bfloat16
        )
        q_te, k_te, v_te = jax.device_put((q_te, k_te, v_te), device)

        te_attention = te_flax.DotProductAttention(
            head_dim=HEAD_DIM,
            num_attention_heads=NUM_HEADS,
            num_gqa_groups=NUM_HEADS,
            attn_mask_type="causal",
            transpose_batch_sequence=False,
        )

        te_vars = te_attention.init(
            jax.random.key(0),
            q_te,
            k_te,
            v_te,
            deterministic=True,
        )

        def te_fn(q, k, v):
            return te_attention.apply(te_vars, q, k, v, deterministic=True)

        row = {"seq_len": sl}

        row["sdpa_xla_bf16_ms"] = benchmark_attention(
            lambda q, k, v: jax.nn.dot_product_attention(
                q, k, v, implementation="xla", is_causal=True
            ),
            q_te, k_te, v_te,
            warmup=2,
            repeats=20,
        )

        if HAS_CUDNN_SDPA:
            row["sdpa_cudnn_bf16_ms"] = benchmark_attention(
                lambda q, k, v: jax.nn.dot_product_attention(
                    q, k, v, implementation="cudnn", is_causal=True
                ),
                q_te, k_te, v_te,
                warmup=2,
                repeats=20,
            )

        row["te_bf16_ms"] = benchmark_attention(
            te_fn,
            q_te, k_te, v_te,
            warmup=2,
            repeats=20,
        )

        te_results.append(row)


    headers = ["Seq len", "SDPA XLA bf16 causal (ms)"]
    if HAS_CUDNN_SDPA:
        headers.append("SDPA cuDNN bf16 causal (ms)")
    headers.append("TE DotProductAttention bf16 causal (ms)")

    table_rows = []
    for r in te_results:
        row = [
            r["seq_len"],
            f"{r['sdpa_xla_bf16_ms']:.2f}",
        ]

        if HAS_CUDNN_SDPA:
            row.append(f"{r['sdpa_cudnn_bf16_ms']:.2f}")

        row.append(f"{r['te_bf16_ms']:.2f}")
        table_rows.append(row)

    show_table(
        headers,
        table_rows,
        title=f"TransformerEngine sequence-length sweep — B={TE_BATCH}, N={NUM_HEADS}, H={HEAD_DIM}",
        aligns=["right"] * len(headers),
    )


    fig, ax = plt.subplots(figsize=(8, 5))
    seq_lens = [r["seq_len"] for r in te_results]

    ax.plot(
        seq_lens,
        [r["sdpa_xla_bf16_ms"] for r in te_results],
        "d-",
        label="SDPA XLA bf16 causal",
        color="#8250df",
    )

    if HAS_CUDNN_SDPA:
        ax.plot(
            seq_lens,
            [r["sdpa_cudnn_bf16_ms"] for r in te_results],
            "^-",
            label="SDPA cuDNN bf16 causal",
            color="#1a7f37",
        )

    ax.plot(
        seq_lens,
        [r["te_bf16_ms"] for r in te_results],
        "o-",
        label="TE DotProductAttention bf16 causal",
        color="#d1242f",
    )

    ax.set_xlabel("Sequence length")
    ax.set_ylabel("Median ms per call")
    ax.set_title("Causal attention latency vs sequence length")
    ax.legend()
    ax.grid(True, alpha=0.25)
    ax.set_xticks(seq_lens)

    fig.tight_layout()
    plt.show()

১২৮ থেকে ২০৪৮ পর্যন্ত প্রতিটি সিকোয়েন্স দৈর্ঘ্যের জন্য আপনার একটি টেবিল সারি এবং একটি প্লট পয়েন্ট পাওয়া উচিত। এই bf16 বেঞ্চমার্কটি একই কজাল-অ্যাটেনশন ওয়ার্কলোডে TransformerEngine-কে JAX SDPA-এর সাথে তুলনা করে, কিন্তু TransformerEngine-এর সম্পূর্ণ পারফরম্যান্সের সম্ভাবনা সাধারণত Hopper এবং Blackwell GPU-গুলিতে দেখা যায় যখন FP8 অটোকাস্ট উপলব্ধ থাকে।

FP8 পথটি পরিদর্শন করুন

হপার জিপিইউ-তে (কম্পিউট ক্যাপাবিলিটি >= ৯.০, যেমন H100), অতিরিক্ত থ্রুপুটের জন্য TransformerEngine FP8-এ অ্যাটেনশন চালাতে পারে। FP8 ডাইনামিক স্কেলিং ফ্যাক্টর গণনা করার জন্য DelayedScaling রেসিপি ব্যবহার করে, যা প্রতিটি টেনসরের অ্যাবসোলিউট-ম্যাক্স হিস্ট্রি ট্র্যাক করে।

  • ফরোয়ার্ড পাসের জন্য E4M3 ফরম্যাট (৪ এক্সপোনেন্ট, ৩ ম্যান্টিসা বিট)
  • ব্যাকওয়ার্ড পাসের জন্য E5M2 ফরম্যাট (৫ এক্সপোনেন্ট, ২ ম্যান্টিসা বিট)

যদি GPU FP8 সমর্থন না করে, তাহলে কোডটি না চালালে কেমন দেখাবে তা এই সেলে দেখানো হয়েছে।

if HAS_TE:
    from transformer_engine.common.recipe import DelayedScaling, Format

    gpu_name = f"{device} {getattr(device, 'device_kind', '')}".lower()
    HAS_FP8 = any(
        tag in gpu_name
        for tag in ["h100", "h200", "b100", "b200", "gb200", "blackwell"]
    )

    fp8_recipe = DelayedScaling(
        margin=0,
        fp8_format=Format.HYBRID,
        amax_history_len=1024,
        amax_compute_algo="max",
    )

    if HAS_FP8:
        FP8_SEQ_LEN = 2048
        FP8_BATCH = BATCH

        rk = jax.random.key(9000)
        rk1, rk2, rk3 = jax.random.split(rk, 3)

        q_fp8 = jax.random.normal(
            rk1, (FP8_BATCH, FP8_SEQ_LEN, NUM_HEADS, HEAD_DIM), dtype=jnp.bfloat16
        )
        k_fp8 = jax.random.normal(
            rk2, (FP8_BATCH, FP8_SEQ_LEN, NUM_HEADS, HEAD_DIM), dtype=jnp.bfloat16
        )
        v_fp8 = jax.random.normal(
            rk3, (FP8_BATCH, FP8_SEQ_LEN, NUM_HEADS, HEAD_DIM), dtype=jnp.bfloat16
        )
        q_fp8, k_fp8, v_fp8 = jax.device_put((q_fp8, k_fp8, v_fp8), device)

        fp8_attention = te_flax.DotProductAttention(
            head_dim=HEAD_DIM,
            num_attention_heads=NUM_HEADS,
            num_gqa_groups=NUM_HEADS,
            attn_mask_type="causal",
            transpose_batch_sequence=False,
        )

        bf16_vars = fp8_attention.init(
            jax.random.key(0),
            q_fp8,
            k_fp8,
            v_fp8,
            deterministic=True,
        )

        bf16_out = block_tree(
            fp8_attention.apply(
                bf16_vars,
                q_fp8,
                k_fp8,
                v_fp8,
                deterministic=True,
            )
        )

        with te.autocast(enabled=True, recipe=fp8_recipe):
            fp8_vars = fp8_attention.init(
                jax.random.key(1),
                q_fp8,
                k_fp8,
                v_fp8,
                deterministic=True,
            )
            fp8_out = block_tree(
                fp8_attention.apply(
                    fp8_vars,
                    q_fp8,
                    k_fp8,
                    v_fp8,
                    deterministic=True,
                )
            )

        max_diff_fp8 = float(jnp.max(jnp.abs(
            bf16_out.astype(jnp.float32) - fp8_out.astype(jnp.float32)
        )))

        show_table(
            ["", "Value"],
            [
                ("GPU", getattr(device, "device_kind", str(device))),
                ("Input dtype", str(q_fp8.dtype)),
                ("bf16 output dtype", str(bf16_out.dtype)),
                ("FP8 autocast output dtype", str(fp8_out.dtype)),
                ("Output shape", str(fp8_out.shape)),
                ("Max |TE bf16 - TE FP8 autocast|", f"{max_diff_fp8:.2e}"),
            ],
            title="FP8 attention with TransformerEngine",
        )

    else:
        show_table(
            ["", "Value"],
            [
                ("GPU", getattr(device, "device_kind", str(device))),
                ("FP8 support", "No detected support; requires Hopper/Blackwell-class GPU"),
            ],
            title="FP8 attention — not available on this GPU",
        )

        print()
        print("The FP8 path uses TransformerEngine autocast:")
        print()
        print("  with te.autocast(enabled=True, recipe=fp8_recipe):")
        print("      out = fp8_attention.apply(vars, q, k, v, deterministic=True)")

else:
    print("TransformerEngine not available — FP8 section skipped.")

L4-এ আপনি একটি টেবিল দেখতে পাবেন, যেখানে আপনার GPU-এর নাম উল্লেখ থাকবে এবং কোনো FP8 সাপোর্ট শনাক্ত হয়নি বলে জানানো হবে। এর পরেই দুটি প্রিন্টেড লাইন থাকবে, যেখানে একটি Hopper GPU-তে ব্যবহৃত te.autocast কলটি দেখানো হবে। ওই কোড স্নিপেটটি সংরক্ষণ করুন: কল সাইটে FP8-এর জন্য শুধু এই একটি পরিবর্তনই প্রয়োজন।

১১. পরিষ্কার করুন

লোডব্যালেন্সার এবং পার্সিস্টেন্ট ভলিউম সহ জুপিটার ওয়ার্কলোডটি মুছে ফেলুন:

kubectl delete -f deploy/jupyter.yaml

ক্লাস্টার, নোড পুল, ভিপিসি এবং সার্ভিস অ্যাকাউন্ট ধ্বংস করুন:

cd terraform
terraform destroy

নির্দেশিত হলে yes টাইপ করুন, তারপর নিশ্চিত করুন যে পিছনে কিছু ফেলে রাখা হয়নি:

gcloud container clusters list
gcloud compute instances list

এই প্রজেক্টের জন্য উভয়ই খালি থাকা উচিত। যদি আপনি শুধু এই সিরিজের জন্য একটি প্রজেক্ট তৈরি করে থাকেন, তাহলে আপনি এর পরিবর্তে ক্লাউড কনসোল থেকে পুরো প্রজেক্টটি মুছে ফেলতে পারেন।

১২. অভিনন্দন

আপনি হাতে লেখা অ্যাটেনশন কম্পিউটেশন থেকে জিপিইউ-অপ্টিমাইজড ফিউজড কার্নেলে স্থানান্তরিত হয়েছেন এবং বাস্তব হার্ডওয়্যারে এর পার্থক্য পরিমাপ করেছেন।

আপনি যা শিখেছেন

  • নেইভ অ্যাটেনশন কাজ করে, কিন্তু এটি একাধিক জিপিইউ কার্নেল চালু করে এবং সম্পূর্ণ অ্যাটেনশন ম্যাট্রিক্সটিকে মেমরিতে বাস্তবায়িত করে।
  • jax.nn.dot_product_attention গণনাকে একটি একক অপারেশনে একীভূত করে, এবং implementation=None থাকলে JAX স্বয়ংক্রিয়ভাবে সেরা ব্যাকএন্ডটি বেছে নেয়।
  • implementation="cudnn" বিকল্পটি এনভিডিয়ার cuDNN ফিউজড অ্যাটেনশন কার্নেল ব্যবহারে বাধ্য করে, যা দীর্ঘ সিকোয়েন্সের ক্ষেত্রে দ্রুততম, এবং এর জন্য bfloat16 বা float16 ইনপুট ও compute capability 8.0 বা তার নতুন সংস্করণ প্রয়োজন।
  • is_causal=True সহ কার্যকারণ মাস্কিং ফিউজড কার্নেলের অন্তর্নির্মিত — কোনো ম্যানুয়াল মাস্ক ম্যাট্রিক্সের প্রয়োজন নেই।
  • GQA এবং MQA ইনফারেন্সের সময় মেমরি সাশ্রয়ের জন্য KV হেড কমিয়ে দেয়, এবং dot_product_attention স্বয়ংক্রিয়ভাবে ব্রডকাস্টিং পরিচালনা করে।
  • TransformerEngine, Hopper GPU-গুলিতে ঐচ্ছিক FP8 প্রিসিশন সহ ফিউজড অ্যাটেনশন প্রদান করে।

পরবর্তী পদক্ষেপ

  • কোডল্যাব ৬: একাধিক জিপিইউ জুড়ে JAX ট্রেনিং স্কেল করা আপনাকে উভয় L4 জুড়ে অ্যারে শার্ডিং করা এবং সমান্তরালভাবে ট্রেনিং স্টেপগুলো চালানোর পদ্ধতি দেখাবে।
  • HEAD_DIM পরিবর্তন করুন (৩২, ৬৪, ১২৮ চেষ্টা করে দেখুন) এবং লক্ষ্য করুন cuDNN কোন মানগুলো গ্রহণ করে ও টাইমিং-এর কী পরিবর্তন হয়।
  • SEQ_LEN বাড়ান (২৫৬, ৫১২, ১০২৪, ২০৪৮ চেষ্টা করে দেখুন) এবং মেমরি ব্যবহার ও cuDNN-এর গতিবৃদ্ধির অনুপাত পর্যবেক্ষণ করুন।
  • ৮টি কোয়েরি হেডের বিপরীতে K=2 এবং K=1 কেভি হেড ব্যবহার করে MHA, GQA, এবং MQA তুলনাটি পুনরায় চালান, এবং ইনফারেন্সের জন্য কেভি ক্যাশে সাশ্রয়ের কারণ ব্যাখ্যা করুন।

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