کاربرد اولیه
تابعی شامل توابع Tokamax را در نظر بگیرید که روی یک پردازنده گرافیکی H100 اجرا میشود:
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 مجاز است بهترین پیادهسازی را برای هر شکل هسته انتخاب کند. حتی مجاز است پیادهسازیهای مختلفی را برای forward pass و gradient انتخاب کند. همچنین همیشه پشتیبانی میشود، زیرا میتواند به یک پیادهسازی XLA implementation='xla' بازگردد.
با این حال، ممکن است بخواهید یک پیادهسازی خاص از هسته را انتخاب کنید و در صورت عدم پشتیبانی، با شکست مواجه شوید. برای مثال، implementation="mosaic" در صورت امکان سعی میکند از یک هسته پردازنده گرافیکی Pallas:Mosaic استفاده کند و اگر به هر دلیلی پشتیبانی نشود، یک استثنا ایجاد میکند. برای مثال، ورودیهای FP64 پشتیبانی نمیشوند یا از پردازندههای گرافیکی قدیمیتر استفاده میکنند.
گرادیان را ارزیابی کنید
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)
اتوتیونینگ
برای به دست آوردن بهترین عملکرد، تمام هستههای Tokamax را در f_grad به صورت خودکار تنظیم کنید:
autotune_result: tokamax.AutotuningResult = tokamax.autotune(f, x, scale)
autotune_result میتواند به عنوان یک مدیر زمینه، با استفاده از پیکربندیهای تنظیمشده خودکار برای همه هستههای Tokamax در f_grad استفاده شود:
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.autotune با ارثبری از کلاس tokamax.Op و بازنویسی متد tokamax.Op._get_autotuning_configs برای تعریف فضای جستجوی autotuning، به صورت خودکار تنظیم کنند.
توجه داشته باشید که تنظیم خودکار اساساً غیرقطعی است: اندازهگیری زمان اجرای هسته پر سر و صدا است. از آنجایی که پیکربندیهای مختلف انتخاب شده در طول تنظیم خودکار میتوانند به اعداد مختلف منجر شوند، این یک منبع بالقوه برای عدم قطعیت عددی است. سریالسازی و استفاده مجدد از نتایج ثابت تنظیم خودکار، راهی برای اطمینان از اعداد یکسان در طول جلسات است.
سریالسازی
هستهها را میتوان به 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),
)
توجه داشته باشید که توابع سریالیزه شده با هستههای Tokamax، استقلال دستگاه StableHLO استاندارد را از دست میدهند. Tokamax دو تضمین سریالیزه شدن را ارائه میدهد:
- یک تابع deserialized که روی یک دستگاه خاص serial شده است، تضمین میشود که دقیقاً روی همان دستگاهی که برای آن serial شده است، اجرا شود.
- توکامکس همان تضمینهای سازگاری 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) .