অটোটিউনিং এবং হিউরিস্টিকস

কার্নেলের পারফরম্যান্স ইমপ্লিমেন্টেশন-নির্দিষ্ট হাইপারপ্যারামিটার (যেমন, টাইল সাইজ, লেআউট ইত্যাদি) দ্বারা উল্লেখযোগ্যভাবে প্রভাবিত হতে পারে, যা টোকাম্যাক্সে অপ কনফিগ (Op Configs) নামে পরিচিত। অটোটিউনিং- এর মাধ্যমে, প্রতিটি প্রাসঙ্গিক ইনপুটের জন্য এই কনফিগগুলোর সর্বোত্তম মান পরীক্ষামূলকভাবে নির্ধারণ করতে হয়। টোকাম্যাক্স অটোটিউনিং সম্পাদন করার জন্য এবং এই প্রক্রিয়া ও এর আউটপুটগুলোর ব্যবস্থাপনা সহজ করার জন্য একটি ফ্রেমওয়ার্ক প্রদান করে।

উচ্চ-স্তরের সংক্ষিপ্ত বিবরণ

প্রতিটি tokamax.Op ইমপ্লিমেন্টেশন তার টিউনযোগ্য হাইপারপ্যারামিটারগুলোর সেট এবং প্রতিটি কনফিগারেশনের সম্ভাব্য মানের পরিসর tokamax.Op._get_autotuning_configs ( উদাহরণ ) ফাইলে সংজ্ঞায়িত করে। অটোটিউনারটি সেরা পারফরম্যান্স (সর্বনিম্ন এক্সিকিউশন টাইম) সম্পন্ন কনফিগারেশনটি শনাক্ত করার জন্য এই পরিসর জুড়ে একটি পুঙ্খানুপুঙ্খ অনুসন্ধান চালায়। এর ফলাফল বিভিন্ন উপায়ে ব্যবহার করা যেতে পারে: সরাসরি আপনার প্রোগ্রামে একটি কনটেক্সট ম্যানেজার হিসেবে, সিরিয়ালাইজ করে পরবর্তী ব্যবহারের জন্য আপনার ব্যক্তিগত ক্যাশে সংরক্ষণ করে, অথবা গ্লোবাল লাইব্রেরি-ব্যাপী ক্যাশের অংশ হিসেবে।

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

অটোটিউনার ইনপুট ফরম্যাট

কলযোগ্য ফাংশন

এক বা একাধিক টোকাম্যাক্স অপস সহ কোনো কলযোগ্য ফাংশন দেওয়া হলে, অটোটিউনারটি 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)

Tokamax VJP অপস তাদের নিজস্ব Op ক্লাস হিসেবে প্রয়োগ করা যেতে পারে এবং তাদের নিজস্ব সার্চ স্পেস ও অনন্য _get_autotuning_configs. যেহেতু VJP অপস-এর অটোটিউনিং-এর জন্য সরাসরি কল করার মতো কোনো API নেই, তাই আপনি আপনার ফরওয়ার্ড ফাংশন ব্যবহার করে jax.grad কল করে Tokamax VJP অপস-এর পুলব্যাক ফাংশনটি পেতে পারেন এবং তারপর সেই পুলব্যাক ফাংশনটি অটোটিউনারে পাস করতে পারেন।

# 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 Op) কল করা হয়, তখন অপ-টি প্রথমে আর্গুমেন্টগুলোর সাথে "বাইন্ড" হয়। এই প্রক্রিয়াটি ইনপুটগুলোকে ক্যানোনিকালাইজ ও ভ্যালিডেট করে এবং অপ-টি ও তার ইনপুট আর্গুমেন্টগুলো নিয়ে একটি 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 স্ক্রিপ্টটি প্রোফাইলিং রানের সময় XProf দ্বারা তৈরি একটি xplane প্রোটো ইনপুট হিসেবে গ্রহণ করে এবং BoundArguments এর একটি JSON আউটপুট দেয়। অটোটিউনার এই JSON ফাইলটি পড়ে প্রোফাইলে থাকা সমস্ত Tokamax Ops-কে অটোটিউন করতে পারে।

ধরা যাক, প্রোফাইলিং থেকে আপনি my_model.xplane.pb নামের একটি এক্সপ্লেন পেয়েছেন, আপনি এতে ব্যবহৃত সমস্ত টোকাম্যাক্স অপস (Tokamax ops) এইভাবে এক্সট্র্যাক্ট এবং অটোটিউন করতে পারেন:

 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)

অটোটিউনার আউটপুট ব্যবহার

অটোটিউনিং প্রক্রিয়াটি সর্বদা একটি 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 সিরিয়ালাইজ করে পুনরায় ব্যবহার করতে চাইতে পারেন: (১) সার্চ স্পেসের আকারের উপর নির্ভর করে অটোটিউনিং কম্পিউট টাইমের দিক থেকে ব্যয়বহুল হতে পারে, (২) কার্নেল এক্সিকিউশন টাইম নয়েজি হতে পারে, যার ফলে বিভিন্ন রানে ভিন্ন ভিন্ন “সেরা” পছন্দ আসতে পারে এবং এর ফলে নিউমেরিক্যাল নন-ডিটারমিনিজম দেখা দিতে পারে।

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)

load ব্যবহার করে ফলাফলগুলোকে একটি JSON স্ট্রিং হিসেবে পড়ে একাধিক `AutotuningResult` অবজেক্টকে একটি একক অবজেক্টে একত্রিত করাও সম্ভব।

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:
  ...