پرش به محتوای اصلی
پرش به محتوای مقاله

Flash-MSA هزینه آموزش مدل‌های میلیون‌توکنی را با هسته‌های پراکنده کاهش داد

·۲۲ تیر ۱۴۰۵۹ دقیقه مطالعه۱ بازدید
هسته‌های توجه پراکنده برای تسریع آموزش میلیون‌توکن با Flash-MSA
هسته‌های توجه پراکنده برای تسریع آموزش میلیون‌توکن با Flash-MSA
اشتراک‌گذاری
واقعاً چه چیز جدید است؟

نخستین پیاده‌سازی متن‌باز و بهینه از توجه پراکنده MSA برای معماری GQA که گلوگاه حافظه در آموزش مدل‌های میلیون‌توکنی را روی GPUهای جدید انویدیا حل می‌کند.

اگر قصد دارید مدلی را با پنجره متنی یک میلیون توکن آموزش دهید، احتمالاً با کرش‌های پی‌درپی حافظه GPU مواجه شده‌اید. دلیل این اتفاق، رشد درجه-دوم (quadratic) حافظه است که معمولاً هنگام آموزش مدل‌هایی با پنجره‌های متنی عظیم، باعث اتمام فضای حافظه گرافیکی می‌شود. برای حل این بحران، در ۱۲ ژوئیه ۲۰۲۶، توسعه‌دهنده‌ای ابزار Flash-MSA را عرضه کرد: نخستین هسته‌های آموزشی متن‌باز و بهینه برای توجه پراکنده Minimax (MSA) که به‌طور اختصاصی برای معماری‌های NVIDIA Hopper (H100) و Blackwell (B200) طراحی شده‌اند.

بیشتر مدل‌های پیشرو برای افزایش سرعت استنتاج (Inference) — لحظه‌ای که مدل واقعاً جواب تولید می‌کند، شبیه خودِ آشپزی و نه دوره‌ی آموزش آشپز — از توجه پراکنده استفاده می‌کنند، اما کدهای بهینه برای آموزش این مدل‌ها تا کنون یک راز تجاری بوده‌اند. این رویکرد بهینه‌سازی استنتاج مشابه راهکاری است که دیپ‌سیک برای کاهش تأخیر در تولید پاسخ‌های مدل V4 خود به‌کار گرفت. طبق مستندات پروژه در nanduruganesh.github.io، آزمایشگاه‌های غربی در پیاده‌سازی فرمول‌های توجه پراکنده (مانند مدل‌های GLM-5.2 یا DSv4) ناکام بودند، زیرا این مدل‌ها بر پایه توجه نهانی چندسر (MLA) بنا شده‌اند که خارج از آزمایشگاه‌های خاص چینی نادر است. Flash-MSA این مشکل را با تطبیق توجه پراکنده با توجه پرس‌وجوی گروهی (GQA) حل کرده است.

همان‌طور که در تحلیل‌های پیشین ما درباره‌ی بهینه‌سازی حافظه در مدل‌های زبانی اشاره کردیم، مدیریت حجم داده‌ها در پنجره‌های متنی بلند، سخت‌ترین چالش سخت‌افزاری است.

مکانیسم پراکندگی بلوکی

برخلاف توجه پراکنده DeepSeek (DSA) که جفت‌های کلید-مقدار (KV) را تک‌تک انتخاب می‌کند، MSA از پراکندگی بلوکی استفاده می‌کند. در این روش، جفت‌های KV در بلوک‌های ۱۲۸تایی و از طریق max-pooling روی امتیازات پروکسی انتخاب می‌شوند.

معماری Flash-MSA با هسته‌های توجه پراکنده، آموزش متن میلیون‌توکنی را تسریع می‌کند.

این تغییر باعث بهبود ویژگی‌های کشینگ در هسته‌ها می‌شود. چون سیستم به‌جای توکن‌های تک‌تک، فقط اندیس‌های بلوکی را ذخیره می‌کند، تنها بخش «گذر پیشرو پروکسی» در مرحله آموزش است که نسبت به طول متن رابطه‌ای درجه-دوم دارد و سایر بخش‌ها با استفاده از این بلوک‌های پراکنده و کش‌شده اجرا می‌شوند.

طراحی و اجرای هسته

برای اجرای بهینه MSA، هسته باید از فشار بیش از حد به ثبات‌ها (Registers) و حافظه مشترک جلوگیری کند. ترتیب اجرا به این شکل است: ابتدا توجه پروکسی، سپس توجه اصلی پراکنده و در نهایت ارسال خروجی به لایه بعد همراه با ذخیره Log-Sum-Exp (LSE) برای مرحله پس‌انتشار.

هسته‌های توجه پراکنده برای تسریع آموزش میلیون‌توکن با Flash-MSA

  • توجه پروکسی: در این مرحله خروجی جمع نمی‌شود، بلکه فقط امتیازات و اندیس‌های k-برتر ردیابی می‌شوند. برای کاهش مصرف حافظه، توسعه‌دهنده از یک مرتب‌سازی درجی (Insertion Sort) برای مقادیر در ثبات‌ها استفاده کرده و بلوک‌های کلید را نصف کرده است.
  • توجه اصلی: این یک گذر پیشرو از توجه برق‌آسا (Flash Attention) با پراکندگی بلوکی است که با استفاده از ترفندی در MoBA، مانند یک توجه برق‌آسا با طول متغیر عمل می‌کند.
  • گذر پس‌انتشار ادغام‌شده: مرحله پس‌انتشار (Backpropagation) — فرآیند اصلاح خطای مدل از انتها به ابتدا برای یادگیری بهتر — در اینجا با توجه پروکسی و اصلی ادغام شده است. این کار لازم است چون سیگنال‌های آموزشی پروکسی به هر دو توزیع احتمالی به‌طور هم‌زمان نیاز دارند. با معکوس کردن نگاشت اندیس‌های بلوکی، هسته تکه‌های پرس‌وجو (Query Chunks) را به‌گونه‌ای زمان‌بندی می‌کند که استفاده مجدد از بلوک‌های پراکنده KV در حافظه مشترک به حداکثر برسد.

رفع گلوگاه واگرایی KL

محاسبه زیان واگرایی KL برای سر پروکسی معمولاً به ایجاد توزیع‌های احتمالی کامل نیاز دارد که باعث کندی شدید آموزش به دلیل عملیات خواندن/نوشتن حجیم در حافظه مشترک می‌شود.

به نقل از مستندات فنی، هسته Flash-MSA از یک ترفند ریاضی برای پس‌انتشار اتمی استفاده می‌کند. توسعه‌دهنده با بسط دادن ترم KL دریافت که گرادیان مربوط به امتیاز پروکسی صرفاً تفاضل احتمال پروکسی و احتمال اصلی ($p_{px,i} - p_i$) است. این یعنی هسته می‌تواند گرادیان‌ها را بدون نیاز به تجسم کامل توزیع KL محاسبه کند و از عملیات سنگین حافظه اجتناب کند.

اعتبارسنجی و عملکرد

صحت این هسته با مقایسه آن با یک پیاده‌سازی eager در PyTorch تأیید شده است. در بررسی‌های مختلف (اندازه دسته ۱ تا ۴ و طول توالی تا ۸٬۱۹۲)، شباهت کسینوسی (Cosine Similarity) برای خروجی‌های پیشرو و گرادیان‌های پس‌انتشار به‌طور مداوم بالا بود. برای مثال، در تست با دسته ۴ و طول ۸٬۱۹۲، گرادیان‌های تصویر پیشرو به شباهت کسینوسی ۰.۹۹۹۵ رسیدند که کاملاً در محدوده تحمل دقت bf16 (یعنی ۰.۰۱) قرار دارد.

معماری Flash-MSA با هسته‌های توجه پراکنده، آموزش متن میلیون‌توکنی را تسریع می‌کند.

محدودیت‌های فعلی و مقیاس‌پذیری آینده

با وجود این پیشرفت‌ها، گذر پس‌انتشار ادغام‌شده در حال حاضر با مشکل اشغال پایین (Low Occupancy) مواجه است. نیاز شدید به ثبات‌ها (۱۳۸ ثبات برای هر رشته) باعث می‌شود سیستم تنها به یک CTA در هر SM محدود شود و اشغال نظری را به ۱۲.۵٪ برساند (در حالی که در Flash-Attention استاندارد این عدد ۱۸.۷۵٪ است).

برای مقیاس‌پذیری بیشتر، مسیرهای زیر در دست بررسی هستند:

  • موازی‌سازی زمینه (CP): پیاده‌سازی all-gather در سطح سر (Head-wise) یا موازی‌سازی حلقوی (Ring-style) برای جلوگیری از سرریز حافظه در طول‌های بسیار زیاد.
  • بهینه‌سازی مسیریاب: استفاده از IndexShare برای به اشتراک‌گذاری سرهای پروکسی بین لایه‌ها؛ تکنیکی که پیش‌تر در مدل GLM پایدار بودن آن اثبات شده است.
  • آموزش با دقت پایین: آموزش اندیس‌ساز با دقت پایین برای تطبیق با رفتار استنتاج و افزایش سرعت کلی عملیات.

این پیاده‌سازی مانع بزرگی را برای توسعه‌دهندگانی که می‌خواهند مدل‌های با زمینه بلند را بدون نیاز به معماری‌های پیچیده MLA آموزش دهند، برطرف می‌کند. با در دسترس قرار دادن MSA برای مدل‌های مبتنی بر GQA، ابزار Flash-MSA راهی را برای مدل‌های با وزن‌های باز (Open-weights) باز می‌کند تا بتوانند با قابلیت‌های زمینه-بلند مدل‌های پیشرو رقابت کنند.

گام بعدی شما

  • اگر روی مدل‌های GQA کار می‌کنید، مستندات Flash-MSA را برای جایگزینی مکانیزم توجه بررسی کنید.
  • تست‌های مقایسه‌ای خود را با استفاده از معیار شباهت کسینوسی برای تأیید صحت گرادیان‌ها انجام دهید.
  • برای کاهش هزینه‌های حافظه، استراتژی‌های پراکندگی بلوکی را در لایه‌های میانی مدل خود آزمایش کنید.

اما داستان سخت‌افزاری این تحول حتی شگفت‌انگیزتر است — به تحلیل ما درباره‌ی تراشه‌های Blackwell مراجعه کنید.

چرا این موضوع مهم است؟

این ابزار با کاهش پیچیدگی حافظه از درجه دوم به تقریباً خطی، هزینه و زمان آموزش مدل‌های با پنجره متنی عظیم را به‌شدت کاهش می‌دهد. این تحول بر اساس تخصص در معماری GPUهای Hopper و Blackwell، امکان رقابت مدل‌های کوچک‌تر را با مدل‌های پیشرو در پردازش متون طولانی فراهم می‌کند.

تأثیر برای ایران

این ابزار برای پژوهشگران مدل‌های بنیادی در ایران که با محدودیت منابع سخت‌افزاری (GPU) روبرو هستند، امکان آموزش بهینه‌تر مدل‌های Long-context را فراهم می‌کند.

·نگاه ما
تحریریه دات‌هوش

دسترسی به کدهای بهینه آموزش (Training Kernels) همیشه نقطه تمایز آزمایشگاه‌های غول‌پیکر با جامعه متن‌باز بوده است. Flash-MSA با انتقال مکانیزم توجه پراکنده از MLA به GQA، در واقع دموکراتیزه کردن آموزش مدل‌های Long-context را ممکن می‌کند. این یعنی مدل‌های وزن‌باز دیگر مجبور نیستند برای رسیدن به پنجره‌های متنی میلیونی، منتظر انتشار مدل‌های بسته بمانند یا از معماری‌های بسیار خاص چینی تقلید کنند.

منابع

این گزارش با خط‌لولهٔ خودکار دات‌هوش از منابع معتبر جهانی تدوین و زیر نظر تحریریه منتشر شده است. روش کار ما

گفتگو

پنج‌شنبه‌های هوش‌محور

بسته‌ی هفتگی دات‌هوش

۵ خبر، ۲ ابزار، ۱ پرامپت در هر شماره. به‌زودی راه‌اندازی می‌شود — هر پنج‌شنبه صبح.

خبر کلیدی
ابزار کاربردی
پرامپت حرفه‌ای
تحلیل پژوهش
به‌زودی
زاویه‌ی ایرانی
به‌زودی
تمرین این هفته
به‌زودی

راهنماهای دات‌هوش

راهنماهای کاربردیِ دات‌هوش برای کار با هوش مصنوعی — از همین‌جا شروع کنید:

دات‌هوش

راهنمای فارسی هوش مصنوعی — با نگاه به ایران

اخبار روزانه، معرفی ابزارها و مدل‌ها، و آموزشِ کار با هوش مصنوعی؛ همیشه با این پرسش که از ایران چه چیزی کار می‌کند و چه چیزی نه.