মৌলিক ব্যবহার
একটি H100 GPU-তে চলমান Tokamax ফাংশন ধারণকারী একটি ফাংশন বিবেচনা করুন:
import jax
import jax.numpy as jnp
import tokamax
def loss(x, scale):
x = tokamax.layer_norm(
x, scale=scale, offset=None, implementation="triton"
)
x = tokamax.dot_product_attention(x, x, x, implementation="xla_chunked")
x = tokamax.layer_norm(x, scale=scale, offset=None, implementation=None)
x = tokamax.dot_product_attention(x, x, x, implementation="mosaic")
return jnp.sum(x)
f_grad = jax.jit(jax.grad(loss))
implementation=None করা থাকলে, Tokamax প্রতিটি কার্নেল শেপের জন্য সেরা ইমপ্লিমেন্টেশনটি বেছে নিতে পারে। এমনকি ফরোয়ার্ড পাস এবং গ্রেডিয়েন্টের জন্য আলাদা ইমপ্লিমেন্টেশন বেছে নেওয়ারও অনুমতি রয়েছে। এটি সর্বদা সমর্থিত থাকবে, কারণ এটি implementation='xla' মতো একটি XLA ইমপ্লিমেন্টেশনে ফিরে যেতে পারে।
তবে, আপনি কার্নেলের একটি নির্দিষ্ট ইমপ্লিমেন্টেশন বেছে নিতে চাইতে পারেন এবং সেটি অসমর্থিত হলে ব্যর্থ হতে পারেন। উদাহরণস্বরূপ, implementation="mosaic" সম্ভব হলে একটি Pallas:Mosaic GPU কার্নেল ব্যবহার করার চেষ্টা করবে এবং কোনো কারণে এটি অসমর্থিত হলে একটি এক্সেপশন থ্রো করবে। যেমন, FP64 ইনপুট বা পুরোনো GPU ব্যবহার করা অসমর্থিত হতে পারে।
গ্রেডিয়েন্ট মূল্যায়ন করুন
channels, seq_len, batch_size, num_heads = (64, 2048, 32, 16)
scale = jax.random.normal(jax.random.key(0), (channels,), dtype=jnp.float32)
x = jax.random.normal(
jax.random.key(1),
(batch_size, seq_len, num_heads, channels),
dtype=jnp.bfloat16,
)
out = f_grad(x, scale)
অটোটিউনিং
সর্বোত্তম পারফরম্যান্স পেতে, f_grad এ সমস্ত Tokamax কার্নেল অটোটিউন করুন:
autotune_result: tokamax.AutotuningResult = tokamax.autotune(f, x, scale)
autotune_result একটি কনটেক্সট-ম্যানেজার হিসেবে ব্যবহার করা যেতে পারে, যা f_grad এ সমস্ত Tokamax কার্নেলের জন্য অটোটিউন করা কনফিগারেশনগুলো ব্যবহার করে:
with autotune_result:
out_autotuned = f_grad(x, scale)
একটি সম্ভাব্য ব্যয়বহুল tokamax.autotuning কলের ফলাফলকে সিরিয়ালাইজ করে পুনরায় ব্যবহার করতে:
autotune_result_json: str = autotune_result.dumps()
autotune_result = tokamax.AutotuningResult.loads(autotune_result_json)
ব্যবহারকারীরা tokamax.Op ক্লাস থেকে ইনহেরিট করে এবং অটোটিউনিং সার্চ-স্পেস নির্ধারণ করার জন্য tokamax.Op._get_autotuning_configs মেথডটি ওভাররাইড করে tokamax.autotune এর সাহায্যে তাদের নিজস্ব কার্নেল অটোটিউন করতে পারেন।
মনে রাখবেন যে অটোটিউনিং মূলত অনির্দিষ্ট: কার্নেল এক্সিকিউশন টাইম পরিমাপ করলে তাতে নয়েজ বা গোলমাল দেখা যায়। যেহেতু অটোটিউনিং-এর সময় বেছে নেওয়া বিভিন্ন কনফিগারেশনের ফলে ভিন্ন ভিন্ন নিউমেরিক পাওয়া যেতে পারে, তাই এটি নিউমেরিক্যাল নন-ডিটারমিনিজমের একটি সম্ভাব্য উৎস। সেশন জুড়ে একই নিউমেরিক নিশ্চিত করার একটি উপায় হলো নির্দিষ্ট অটোটিউনিং ফলাফলকে সিরিয়ালাইজ করা এবং পুনরায় ব্যবহার করা।
ক্রমিকীকরণ
কার্নেলগুলোকে StableHLO- তে সিরিয়ালাইজ করা যায়। কার্নেল কলগুলো হলো JAX কাস্টম কল, যা ডিফল্টভাবে jax.export এ নিষিদ্ধ থাকে। সমস্ত Tokamax কার্নেল এক্সপোর্ট করার অনুমতি দেওয়ার জন্য tokamax.DISABLE_JAX_EXPORT_CHECKS ব্যবহার করা আবশ্যক ।
from jax import export
f_grad_exported = export.export(f_grad, disabled_checks=tokamax.DISABLE_JAX_EXPORT_CHECKS)(
jax.ShapeDtypeStruct(x.shape, x.dtype),
jax.ShapeDtypeStruct(scale.shape, scale.dtype),
)
উল্লেখ্য যে, টোকাম্যাক্স কার্নেল দিয়ে সিরিয়ালাইজ করা ফাংশনগুলো স্ট্যান্ডার্ড StableHLO-এর ডিভাইস-স্বাধীনতা হারায়। টোকাম্যাক্স দুটি সিরিয়ালাইজেশন নিশ্চয়তা প্রদান করে:
- কোনো নির্দিষ্ট ডিভাইসে সিরিয়ালাইজ করা একটি ডিসিরিয়ালাইজড ফাংশন ঠিক সেই ডিভাইসেই চলবে, যার জন্য এটি সিরিয়ালাইজ করা হয়েছিল।
- টোকাম্যাক্স JAX-এর মতোই সামঞ্জস্যতার নিশ্চয়তা দেয়: ৬ মাসের পশ্চাৎ সামঞ্জস্যতা।
বেঞ্চমার্কিং
JAX পাইথন ওভারহেড প্রায়শই প্রকৃত অ্যাক্সিলারেটর কার্নেল এক্সিকিউশন টাইমের চেয়ে অনেক বেশি হয়। এর মানে হলো jax.block_until_ready(f_grad(x, scale)) এর সময় পরিমাপ করার প্রচলিত পদ্ধতিটি কার্যকর হবে না। Tokamax-এর এমন কিছু ইউটিলিটি রয়েছে যা শুধুমাত্র অ্যাক্সিলারেটর এক্সিকিউশন টাইম পরিমাপ করে:
f_std, args = tokamax.benchmarking.standardize_function(f, kwargs={'x': x, 'scale': scale})
run = tokamax.benchmarking.compile_benchmark(f_std, args)
bench: tokamax.benchmarking.BenchmarkData = run(args)
পরিমাপের বিভিন্ন কৌশল রয়েছে: উদাহরণস্বরূপ, GPU-তে CUPTI প্রোফাইলার আছে যা run(args, method='cupti') এর মাধ্যমে নির্দিষ্ট করা যায়। এটি কার্নেলকে ইনস্ট্রুমেন্ট করে এবং সামান্য ওভারহেড যোগ করে। ডিফল্ট run(args, method=None) Tokamax-কে পদ্ধতি বেছে নেওয়ার সুযোগ দেয় এবং এটি TPU ও GPU উভয়ের জন্যই কাজ করে। পুনরাবৃত্তির সংখ্যা run(args, iterations=10) বাড়িয়ে বেঞ্চমার্ক নয়েজ কমানো যায়।