تنظیم خودکار و اکتشافات
عملکرد هسته ممکن است به طور قابل توجهی تحت تأثیر پارامترهای خاص پیادهسازی (مثلاً اندازه کاشیها، چیدمانها و غیره) قرار گیرد که در Tokamax به عنوان Op Configs شناخته میشوند. مقادیر بهینه برای این پیکربندیها باید به صورت تجربی برای هر ورودی مورد نظر، از طریق تنظیم خودکار ، تعیین شوند. Tokamax چارچوبی را برای انجام تنظیم خودکار و سادهسازی مدیریت این فرآیند و خروجیهای آن فراهم میکند.
نمای کلی سطح بالا
هر پیادهسازی tokamax.Op مجموعه پارامترهای قابل تنظیم خود و محدوده مقادیری که هر یک از این پیکربندیها میتوانند داشته باشند را در tokamax.Op._get_autotuning_configs ( مثال ) تعریف میکند. autotuner یک جستجوی جامع در این فضا انجام میدهد تا پیکربندی با بهترین عملکرد (کمترین زمان اجرا) را شناسایی کند. نتیجه ممکن است به روشهای مختلفی مورد استفاده قرار گیرد: مستقیماً در برنامه شما به عنوان یک مدیر زمینه، سریالسازی شده و ذخیره شده در حافظه پنهان خصوصی شما برای استفاده بعدی، یا به عنوان بخشی از حافظه پنهان سراسری کتابخانه.
پیکربندی بهینه برای یک عملیات به مجموعه خاصی از ورودیهای مورد استفاده برای فراخوانی عملیات، از جمله اشکال، dtypeها و غیره بستگی دارد. API اتوتیونینگ این لیست از ورودیها را به روشهای مختلفی میپذیرد و در همه موارد، یک شیء AutotuningResult را خروجی میدهد که شامل بهترین پیکربندیها برای هر عملیات در ورودی است.
قالبهای ورودی اتوتیونر
تابع قابل فراخوانی
با توجه به یک تابع قابل فراخوانی مورد نظر با یک یا چند عملیات Tokamax، تنظیمکننده خودکار از فراداده StableHLO برای استخراج اطلاعات مربوط به هر عملیات و ورودیهای مربوطه آن استفاده میکند و هر یک را به صورت خودکار تنظیم میکند.
# Autotune all Tokamax kernels in f, which can be a JAX function with multiple Tokamax ops and non-Tokamax ops as well. Assume f takes a dictionary of args as input.
autotune_result: tokamax.AutotuningResult = tokamax.autotune(f, **args)
# Best possible result
with autotune_result:
out_autotuned = f(**args)
عملیاتهای VJP توکامکس میتوانند به عنوان کلاس عملیات مخصوص به خود پیادهسازی شوند و فضای جستجوی مخصوص به خود و _get_autotuning_configs. از آنجایی که عملیاتهای VJP رابط برنامهنویسی کاربردی (API) مستقیمی ندارند که بتوان آن را برای تنظیم خودکار فراخوانی کرد، میتوانید تابع forward خود را بردارید و jax.grad را برای دریافت تابع pullback با عملیاتهای VJP توکامکس فراخوانی کنید و سپس تابع pullback را به autotuner ارسال کنید.
# This pullback function contains Tokamax VJP ops.
f_grad = jax.grad(f)
autotune_result = tokamax.autotune(f_grad)
with autotune_result:
out = f_grad()
توالی BoundArguments محدود
وقتی یک عملیات Tokamax با مجموعهای از آرگومانهای ورودی فراخوانی میشود، عملیات ابتدا به آرگومانها "متصل" میشود. این فرآیند ورودیها را استانداردسازی و اعتبارسنجی میکند و یک شیء BoundArgument متشکل از عملیات به همراه آرگومانهای ورودی آن ایجاد میکند. تنظیمکننده خودکار میتواند دنبالهای از BoundArguments به عنوان ورودی دریافت کند و هر یک را به صورت خودکار تنظیم کند.
ragged_dot_ba = tokamax.PallasMosaicTpuRaggedDot.bind(x, y, group)
attention_ba = tokamax.PallasMosaicTpuAttention.bind(q, k, v)
autotune_result: tokamax.AutotuningResult = tokamax.autotune([ragged_dot_ba, attention_ba])
with autotune_result:
...
لیست سریالی از BoundArguments
اسکریپت xplane_to_bound_args.py یک xplane proto که توسط XProf در طول اجرای پروفایلینگ تولید شده است را به عنوان ورودی دریافت کرده و یک JSON از BoundArguments را به عنوان خروجی میدهد. تنظیمکننده خودکار میتواند این فایل JSON را بخواند و تمام عملیات Tokamax که در پروفایل ظاهر میشوند را به صورت خودکار تنظیم کند.
فرض کنید از پروفایلینگ، xplane my_model.xplane.pb دارید، میتوانید تمام عملیات Tokamax مورد استفاده را به این صورت استخراج و تنظیم خودکار کنید:
python tokamax/_src/tools/xplane_to_bound_args.py \
--xplane_file=my_model.xplane.pb \
--output_file=/tmp/bound_args.json
import tokamax
bound_args = tokamax.autotuning.bound_args_from_json("/tmp/bound_args.json")
autotune_result = tokamax.autotune(bound_args)
خروجی اتوتیونر
فرآیند تنظیم خودکار (autotuning) همیشه یک شیء AutotuningResult را برمیگرداند که شامل بهترین پیکربندیها برای هر عملیات در ورودی است. شما میتوانید از آن خروجی به روشهای مختلف متناسب با نیازهای خود استفاده کنید.
مدیر زمینه
شما میتوانید مستقیماً از نتیجهی اتوتیونینگ به عنوان یک مدیریتکنندهی زمینه در کد خود استفاده کنید.
# Autotune all Tokamax kernels in f, which can be a JAX function with multiple Tokamax ops and non-Tokamax ops as well. Assume f takes a dictionary of args as input.
autotune_result: tokamax.AutotuningResult = tokamax.autotune(f, **args)
# Best possible result
with autotune_result:
out_autotuned = f(**args)
AutotuningResult سریالی
شما ممکن است بخواهید autotune_result به یکی از دو دلیل زیر سریالایز و دوباره استفاده کنید (1) تنظیم خودکار میتواند از نظر زمان محاسبه، بسته به اندازه فضای جستجو، پرهزینه باشد (2) زمان اجرای هسته میتواند پر سر و صدا باشد و منجر به انتخابهای "بهترین" متفاوت در اجراهای مختلف شود و ممکن است منجر به عدم قطعیت عددی شود.
autotune_result: tokamax.AutotuningResult = tokamax.autotune(f, **args)
autotune_result_json: str = autotune_result.dumps()
autotune_result = tokamax.AutotuningResult.loads(autotune_result_json)
with autotune_result:
out_autotuned = f(**args)
همچنین میتوان چندین شیء AutotuningResult را با خواندن نتایج به صورت یک رشته JSON با استفاده load ، در یک شیء واحد ادغام کرد.
autotune_result_1 = tokamax.AutotuningResult.load(path_to_file_1)
autotune_result_2 = tokamax.AutotuningResult.load(path_to_file_2)
merged_results = autotune_result_1 | autotune_result_2
with merged_results:
...