ترفند باحال برای سرعت بخشیدن به ماشین حساب‌های هوش مصنوعی: خداحافظ پدینگ در FP8!

Fall Back

احتمالاً اگه سر و کاری با یادگیری ماشین، هوش مصنوعی یا برنامه‌نویسی روی کارت‌های گرافیک داشته باشین، اسم FP8 یا GEMM به گوشتون خورده. بذار یه توضیح کوچیک بدم: FP8 یه نوع دقت عددیه که خیلی سبک‌تر از حالت‌های معمولیه (مثل FP32)، یعنی اعداد رو با ۸ بیت نمایش میده، پس هم سرعت میره بالا هم مصرف حافظه میاد پایین. GEMM هم یعنی ضرب و جمع چندتا ماتریس با هم؛ این تو دنیای هوش مصنوعی خیلی پایه‌ست و تقریباً همه کارها بهش ختم میشه!

حالا اینجا یه مشکلی هست که خیلیا بهش گیر دادن: موقع استفاده از FP8، اگه بخوایم ماتریسا رو تو گروه‌های کوچیک‌تر ضرب کنیم (اصطلاح فنی‌ش رو گذاشتن “Grouped GEMM”)، باید هر گروه رو هی پَدینگ کنیم، یعنی مثلاً اگه گروهی کوچیک‌تر از اندازه‌ی استاندارد باشه، باید چندتا صفر بهش بچسبونیم تا ردیف‌هاش به یه سایز خاص (مثلاً ۱۲۸) برسه. این پدینگ هم رم اضافی می‌خواد و سرعت رو خراب می‌کنه، یعنی کلی منابع به هدر میره.

دانشجوها و متخصص‌های این حوزه اومدن رو کارت‌های گرافیکی Hopper (که این روزا مد روزه واسه ترِین کردن مدل‌های بزرگ هوش مصنوعی)، یه روش توپ به اسم “TMA-Adaptive FP8 Grouped GEMM” پیشنهاد دادن. این روش چیه؟

  • به جای این که همیشه همه رو به یه سایز خاص برسونیم (پدینگ کنیم)، اینا اومدن با یه تکنیک هوشمند سراغش رفتن. با چیزی به اسم TMA descriptor pool کار می‌کنه. Descriptor تو کارت گرافیک یعنی یه جور راهنمای اجرایی برای عملیات حافظه؛ اینجا چندتا descriptor مختلف با اندازه‌هایی که از قبل آماده‌شدن (با استفاده از log2(block_M)، یعنی اندازه‌های مختلف بلوک‌ها رو حساب می‌کنه)، تو یه استخر (pool) نگه می‌داره. بعد موقع اجرا، بر اساس نیاز، هوشمند descriptor مناسب رو انتخاب می‌کنه و باعث میشه عملیات ضرب و جمع ماتریس‌ها خیلی سریع و بدون وقفه انجام بشه.

  • یه نکته فنی دیگه هم هست: باید عملیات حافظه همیشه با سایزهای خاصی هماهنگ باشه (مثل ۱۶ بایت تو حافظه سراسری و ۱۲۸ بایت تو حافظه مشترک). اینا هم تو این کار رعایت شده تا مشکلی پیش نیاد.

نتیجه‌ش چی شده؟ با این کار هم سرعت عملیات بین ۱.۷ تا ۲۰.۴ درصد بیشتر شده، هم حافظه مصرفی تا نزدیک ۲۴ درصد (دقیق‌تر بگم ۲۳.۸ درصد) کمتر شده! همه اینا بدون اینکه دقت محاسبات افت کنه یا نتیجه نهایی خراب شه.

خلاصه‌ش اینه: به جای این که همیشه ماتریسا رو تا یه سایز خاص کش بدیم (که هم کندی میاره، هم رم الکی می‌خوره)، اینا با یه روش هوشمندانه، وابسته به سایز هر گروه، عملیات ضرب و جمع رو انجام میدن. کدشون هم رایگان گذاشتن تو گیت‌هاب؛ پس اگه دوست داری خودت تستش کنی، آدرس اینه: https://github.com/sukoncon/TMA-Adaptive-FP8-Grouped-GEMM

راستی اگه جایی اصطلاحاتی مثل “Dual-phase load-store operations” دیدی، بدون منظورش اینه که اطلاعات تو دو مرحله با روشی خاص، خونده و نوشته میشه تا مطمئن شن هیچ دیتایی از دست نمیره یا جا نمی‌مونه.

در کل، این روش باعث میشه کارای مربوط به آموزش یا اجرای مدل‌های هوش مصنوعی (چه موقع Training چه Inference، یعنی چه موقع یادگیری مدل، چه موقع استفاده از مدل یاد گرفته شده) روی کارت‌های Hopper حسابی سریع‌تر و بهینه‌تر شه؛ بدون کلی قِلِق و پدینگ اضافی. دمشون گرم!

منبع: +