تنظیم خودکار و اکتشافات

عملکرد هسته ممکن است به طور قابل توجهی تحت تأثیر پارامترهای خاص پیاده‌سازی (مثلاً اندازه کاشی‌ها، چیدمان‌ها و غیره) قرار گیرد که در 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:
  ...