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

کاهش ترافیک حافظه در برابر کاهش عملیات ریاضی در الگوریتم FlashAttention

·۲۳ شهریور ۱۴۰۵۳۲ دقیقه مطالعه۱ بازدید
راهنما
یادداشت‌های شخصی درباره درک FlashAttention بخش ۱: نمودار معماری و جریان داده در محاسبه توجه سریع با بهینه‌سازی حافظه مشترک GP
یادداشت‌های شخصی درباره درک FlashAttention بخش ۱: نمودار معماری و جریان داده در محاسبه توجه سریع با بهینه‌سازی حافظه مشترک GP
اشتراک‌گذاری
واقعاً چه چیز جدید است؟

تغییر استراتژی از کاهش عملیات ریاضی (FLOPs) به کاهش ترافیک حافظه (IO)؛ به گونه‌ای که سرعت افزایش می‌یابد اما دقت ریاضی مدل به‌طور کامل حفظ می‌شود.

اگر امروز از مدل‌های زبانی با پنجره‌های متنی بسیار بزرگ استفاده می‌کنید، احتمالاً متوجه شده‌اید که سرعت پاسخ‌دهی با افزایش طول متن به‌شدت افت می‌کند. این گلوگاه نه به دلیل کندی پردازش ریاضی، بلکه به دلیل ترافیک سنگین جابه‌جایی داده‌ها در حافظه است. در حالی که مدل‌های زبانی بزرگ (LLM) مدرن برای مقیاس‌پذیری طراحی شده‌اند، یک تک‌سرِ توجه (Attention Head) هنگام پردازش ۸,۱۹۲ توکن، ماتریسی ایجاد می‌کند که به ۴ گیگابایت حافظه نیاز دارد و همین موضوع به گلوگاهی تبدیل می‌شود که سرعت هر مدلی را کاهش می‌دهد.

توجه برق‌آسا (FlashAttention) با تغییر نگاه از «کاهش محاسبات» به «کاهش جابه‌جایی داده»، این مشکل را حل می‌کند. این مکانیزم سلسله‌مراتب حافظه GPU را به عنوان محدودیت اصلی در نظر می‌گیرد و کاهش حرکت داده‌ها را بر کاهش محاسبات خام اولویت می‌دهد. سازوکار این فناوری در سه عبارت خلاصه می‌شود: کاشی‌بندی (Tiling)، سافت‌مکس آنلاین (Online Softmax) و محاسبه مجدد (Recomputation).

همان‌طور که در تحلیل قبلی ما درباره‌ی مدیریت زمینه در مدل‌های زبانی اشاره کردیم، کارایی مکانیزم توجه تعیین می‌کند که مدل در لحظه چه مقدار از بستر متن را می‌تواند پردازش کند. برای درک عمیق‌تر از نحوه عملکرد این لایه‌ها، می‌توانید تجسم لایه‌های توجه و رمزگشایی از مکانیزم کپی-پیست در مدل‌های زبانی را مطالعه کنید که نشان می‌دهد داده‌ها چگونه در این ساختارها جابه‌جا می‌شوند. بسیاری از توسعه‌دهندگان تصور می‌کنند تنها راه افزایش سرعت، کاهش عملیات اعشاری (FLOPs) است. اما در سخت‌افزارهای مدرن، زمانی که صرف جابه‌جایی داده بین حافظه پهنای‌باند بالا (HBM) — که مثل یک انبار بزرگ اما دور است — و حافظه SRAM — که شبیه به میز کار کوچک اما بسیار سریع است — بیشتر از زمان خودِ محاسبات ریاضی است. این همان «توازن حافظه-محاسبه» (Memory-compute tradeoff) است.

گلوگاه حافظه

پیاده‌سازی‌های استاندارد توجه، «محدود به حافظه» (Memory-bound) هستند. آن‌ها معمولاً یک فرآیند سه مرحله‌ای را دنبال می‌کنند: ابتدا ماتریس امتیازات $S$ را محاسبه می‌کنند، سپس سافت‌مکس را برای به دست آوردن احتمالات $P$ اعمال می‌کنند و در نهایت آن را در مقادیر $V$ ضرب می‌کنند.

در یک پیاده‌سازی ابتدایی (Naive)، این فرآیند به این شکل است:
۱. خواندن $Q$ و $K$ از HBM، محاسبه $S$ و نوشتن $S$ در HBM.
۲. خواندن $S$ از HBM، محاسبه $P$ و نوشتن $P$ در HBM.
۳. خواندن $P$ و $V$ به‌صورت بلوکی از HBM، محاسبه $O$ و نوشتن $O$ در HBM.

هر مقدار میانی — یعنی $S$، $P$ و $O$ — باید یک بار نوشته و دوباره خوانده شود. برای دسته‌ای با ۳۲ سر و طول توالی ۸,۱۹۲ با استفاده از فرمت FP16 یا BF16 (دو بایت برای هر عنصر)، یک تنسور از این دست دقیقاً $32 \times 8192^2 \times 2 = 4,294,967,296$ بایت یا ۴ گیگابایت فضا اشغال می‌کند. جابه‌جایی مکرر این شیء ۴ گیگابایتی در گذرگاه حافظه GPU، فارغ از اینکه هسته‌های تنسور (Tensor Cores) چقدر سریع باشند، باعث افت شدید عملکرد می‌شود.

سلسله‌مراتب حافظه GPU و ورودی-خروجی

واحد پردازش گرافیکی (GPU) دارای یک استخر حافظه یکپارچه نیست، بلکه از سلسله‌مراتبی استفاده می‌کند که در آن سرعت با اندازه رابطه عکس دارد:

  • HBM (حافظه دستگاه): حجیم، خارج از تراشه و نسبتاً کند؛ محل ذخیره وزن‌های مدل، تنسورهای $Q/K/V$ و فعال‌سازها (Activations).
  • SRAM (حافظه مشترک): بسیار کوچک‌تر، روی تراشه و به‌شدت سریع؛ جایی که بازاستفاده از داده‌ها ارزان است.
  • ثبات‌ها (Registers): کوچک‌ترین و محلی‌ترین وضعیت حافظه.

FlashAttention یک الگوریتم «آگاه به IO» است چون جابه‌جایی داده بین این سطوح را به حداقل می‌رساند. این الگوریتم تشخیص می‌دهد که اگرچه محاسبات متراکم همچنان مرتبه دوم $\mathcal{O}(N^2 d)$ است، اما ردپای حافظه (Memory Footprint) نباید چنین باشد. با افزایش بازاستفاده از داده‌ها در حالی که روی تراشه (On-chip) هستند، تعداد دفعات رفت‌وبرگشت به HBM کاهش می‌یابد. یک کرنل زمانی «آگاه به IO» می‌شود که جای‌گذاری و جابه‌جایی داده‌ها بخشی از خودِ الگوریتم باشد، نه چیزی که به توالی عملیات تنسوری مجزا سپرده شود.

مکانیزم کاشی‌بندی (Tiling)

FlashAttention یک رویکرد آگاه به IO را معرفی می‌کند. به‌جای محاسبه کامل ماتریس، محاسبات را به بلوک‌های کوچک یا «کاشی‌هایی» تقسیم می‌کند که به‌طور کامل در SRAM سریع روی تراشه جای می‌گیرند.

برای یک سرِ توجه، فرض کنید $Q \in \mathbb{R}^{N_q \times d}$، $K \in \mathbb{R}^{N_k \times d}$ و $V \in \mathbb{R}^{N_k \times d_v}$ باشند. ماتریس‌ها به بلوک‌هایی (مثلاً $Q_1, Q_2, Q_3$) تقسیم می‌شوند. برای هر کاشی پرس‌وجوی $Q_i$، کرنل کاشی‌های متناظر کلید و مقدار $(K_j, V_j)$ را استریم می‌کند.

در داخل کرنل، فرآیند از این حلقه پیروی می‌کند:
۱. بارگذاری یک کاشی از $Q_i$ و یک کاشی از $K_j, V_j$.
۲. محاسبه کاشی امتیازات محلی $S_{ij} = Q_i K_j^T / \sqrt{d} + B_{ij}$ (که در آن $B$ یک ماسک یا بایاس اختیاری است).
۳. اعمال به‌روزرسانی سافت‌مکس به‌صورت محلی با استفاده از آمارهای جاری.
۴. استفاده فوری از این احتمالات برای انباشت سهم متناظر از $V_j$.
۵. دور انداختن کاشی و حرکت به سمت بلوک $K, V$ بعدی.

پس از پردازش هر کاشی، آن را دور می‌اندازد، به این معنی که ماتریس کامل $N \times N$ هرگز در HBM کند متجلی نمی‌شود. این کار پیچیدگی حافظه وضعیت کمکی را از مرتبه دوم $\mathcal{O}(N^2)$ به مرتبه خطی $\mathcal{O}(N)$ تغییر می‌دهد، اگرچه محاسبات همچنان مرتبه دوم باقی می‌مانند. این یک تمایز حیاتی است: وضعیت کمکی خطی به معنای محاسبات خطی نیست. برای $N = 10,000$، همچنان ۱۰۰ میلیون تعامل Q-K وجود دارد؛ FlashAttention صرفاً نیاز به یک تنسور میانی غول‌پیکر برای ذخیره آن‌ها را از بین می‌برد.

ترفند سافت‌مکس آنلاین

سافت‌مکس (Softmax) — که مثل یک سیستم وزن‌دهی است تا مجموع احتمالات برابر یک شود — به‌طور سنتی برای پردازش جریانی دشوار است، زیرا هر عنصر برای پایداری عددی به مقدار حداکثر (Maximum) کل ردیف وابسته است. یک سافت‌مکس پایدار از $m = \max_j x_j$ برای جلوگیری از سرریز (Overflow) استفاده می‌کند، اما در یک جریان داده، ممکن است بلوک‌های بعدی حاوی مقدار حداکثری بزرگ‌تر از بلوک‌های قبلی باشند.

FlashAttention برای حل این مشکل از یک بازگشت «سافت‌مکس آنلاین» استفاده می‌کند. این روش آمارهای کافی — یعنی حداکثر جاری $m$ و مخرج $\ell$ — را نگه می‌دارد که می‌توانند هنگام تغییر مقدار حداکثر، دوباره مقیاس‌بندی شوند.

اگر وضعیت جاری $(m_{\text{old}}, \ell_{\text{old}})$ باشد و یک بلوک جدید دارای حداکثر $m_b$ باشد، حداکثر جهانی جدید $m_{\text{new}} = \max(m_{\text{old}}, m_b)$ خواهد بود. سهم‌های قبلی با ضریب $\alpha = e^{m_{\text{old}} - m_{\text{new}}}$ بازمقیاس می‌شوند. این بر اساس این اتحاد است: $e^{x_j - m_{\text{old}}} e^{m_{\text{old}} - m_{\text{new}}} = e^{x_j - m_{\text{new}}}$.

بازگشت ادغام بلوکی (Blockwise Merge Recurrence):
برای یک ردیف پرس‌وجو، کرنل مقادیر $(m, \ell, a)$ را نگه می‌دارد که در آن $a$ انباشت‌کننده مقدار نرمال‌نشده است. برای یک بلوک امتیاز جدید $s$ و مقادیر $V_b$:

  • $m' = \max(m, \max(s))$
  • $\alpha = e^{m - m'}$
  • $p = e^{s - m'}$
  • $\ell' = \alpha \ell + \sum p_j$
  • $a' = \alpha a + p^T V_b$

در پایان توالی، خروجی نهایی $O = a / \ell$ است. این به کرنل اجازه می‌دهد تا یک نرمال‌ساز جاری و یک مجموع مقادیر وزن‌دار را بدون نیاز به بازگشت به بلوک‌های قبلی حفظ کند. این بازگشت دلیل جبری است که چرا مرزهای کاشی‌بندی، نتیجه نهایی سافت‌مکس متراکم را تغییر نمی‌دهند.

دقت مطلق در برابر تقریب

برخلاف روش‌های توجه پراکنده (Sparse) یا خطی، FlashAttention یک الگوریتم دقیق است. این روش هیچ جفت پرس‌وجو-کلیدی را حذف نمی‌کند، تقریب‌های کم‌رتبه (Low-rank) معرفی نمی‌کند و تابع سافت‌مکس را با یک فرمول کرنلی جایگزین نمی‌کند. این الگوریتم همان تابع ریاضی توجه استاندارد را محاسبه می‌کند: $O = \mathrm{softmax}(QK^T / \sqrt{d} + B) V$.

یک تمایز حیاتی بین «دقت ریاضی» و «همانی بیتی» (Bitwise identity) وجود دارد. چون کاشی‌بندی ترتیب کاهش عملیات اعشاری را تغییر می‌دهد، خروجی ممکن است به دلیل اثرات گرد کردن (Rounding effects) بسیار ناچیز با پیاده‌سازی ابتدایی متفاوت باشد. به همین دلیل است که بک‌اندهای scaled-dot-product attention (SDPA) در PyTorch ممکن است نتایج کمی متفاوت تولید کنند. «دقیق» بودن به تابع ریاضی اشاره دارد، نه بازتولید بیتی دقیق.

آموزش و گذر پس‌رو

در زمان آموزش، سیستم‌های Autograd ابتدایی ماتریس احتمالات کامل $N \times N$ (یعنی $P$) را در HBM ذخیره می‌کنند تا گرادیان‌ها را در گذر پس‌رو (Backward Pass) محاسبه کنند. این کار حافظه عظیمی را مصرف می‌کند.

FlashAttention با ذخیره تنها آمارهای فشرده نرمال‌سازی ردیفی (مانند log-sum-exp) از این کار اجتناب می‌کند. در گذر پس‌رو، کرنل کاشی‌های احتمالات را به‌صورت لحظه‌ای بازسازی می‌کند:
۱. محاسبه مجدد $QK^T$ برای آن کاشی.
۲. بازسازی مقادیر احتمالات محلی با استفاده از آمارهای ذخیره شده.
۳. محاسبه گرادیان‌ها برای $dQ, dK$ و $dV$ با استفاده از اتحادهایی مانند $dV \mathrel{+}= P^T dO$ و $dP = dO \cdot V^T$.
۴. دور انداختن کاشی.

این رویکرد، مقدار کمی محاسبات بیشتر (FLOPs بیشتر) را با کاهش شدید ترافیک حافظه معاوضه می‌کند. در GPUهای مدرن، این معامله تقریباً همیشه سودآور است چون ضرب ماتریسی به‌طور استثنایی با هسته‌های تنسور سازگار است، در حالی که ترافیک HBM و همگام‌سازی (Synchronization) به‌نسبت گران هستند.

تکامل سخت‌افزاری

این الگوریتم همگام با تغییرات معماری GPU و شکاف رو به رشد بین توان محاسباتی و پهنای‌باند حافظه تکامل یافته است:

  • FlashAttention-2: بهبود موازی‌سازی در کاشی‌های توالی و کاهش FLOPهای غیر-ضرب ماتریسی. این نسخه تقسیم‌بندی کار را بهینه کرد تا اشغال GPU (GPU Occupancy) بالاتری تضمین شود، با این شناخت که عملیات غیر-ضرب ماتریسی از همان توان عملیاتی GEMMهای هسته تنسور برخوردار نیستند.
  • FlashAttention-3: هدف این نسخه نسل Hopper انویدیا است. از ناهمگامی (Asynchrony) برای هم‌پوشانی جابه‌جایی داده‌ها، ضرب ماتریسی (GEMM) و عملیات سافت‌مکس استفاده می‌کند و از ویژگی‌های سخت‌افزاری تخصصی برای پنهان کردن تأخیر (Latency) بهره می‌برد.
  • FlashAttention-4: بهینه‌شده برای نسل Blackwell. این نسخه به مقیاس‌بندی نامتقارن سخت‌افزاری می‌پردازد، جایی که توان هسته‌های تنسور سریع‌تر از ترافیک حافظه مشترک رشد کرده است و باعث شده توابع نمایی و جابه‌جایی در SRAM نسبتاً گران‌تر شوند.

سازگاری معماری

این مکانیزم نسبت به معماری سر (Head) مدل مستقل است و می‌تواند با ویژگی‌های مختلف ترکیب شود:

  • اشتراک سر (Head Sharing): به‌طور کامل از توجه چندسر (MHA)، توجه تک-پرس‌وجو (MQA) و توجه گروهی (GQA) پشتیبانی می‌کند. در GQA، چندین سر پرس‌وجو یک سر KV واحد را به اشتراک می‌گذارند، به شرطی که تعداد سرهای پرس‌وجو بر تعداد سرهای KV بخش‌پذیر باشد ($H_q \bmod H_{kv} = 0$).
  • ماسک‌گذاری علّی (Causal Masking): توجه خودبازگشتی را با نادیده گرفتن کامل کاشی‌هایی که بالای مرز علّی قرار دارند مدیریت می‌کند. بلوک‌هایی که کاملاً بالای مرز هستند ماسک شده و نیازی به مشارکت ندارند؛ تنها بلوک‌هایی که با قطر ماتریس تلاقی دارند به ماسک‌گذاری در سطح عنصر نیاز دارند.
  • سایر ویژگی‌ها: پیاده‌سازی‌های تجاری از توالی‌های با طول متغیر، توجه پنجره لغزان (SWA) و بایاس‌های سبک ALiBi پشتیبانی می‌کنند.

جزئیات پیاده‌سازی

در چارچوب‌های عملیاتی، FlashAttention در APIهای سطح بالا ادغام شده است. برای مثال، scaled_dot_product_attention در PyTorch، دراپ‌اوت (Dropout) را بر اساس dropout_p اعمال می‌کند. توسعه‌دهندگان باید در زمان ارزیابی (Evaluation) صراحتاً آن را روی ۰.۰ تنظیم کنند تا غیرفعال شود.

همچنین ذکر این نکته مهم است که در حالی که FlashAttention یک کرنل است، الگوی توجه مدل همچنان می‌تواند پراکنده (Sparse) باشد. اگر مدلی از یک پنجره محلی استفاده کند، FlashAttention می‌تواند کرنلی باشد که آن الگو را اجرا می‌کند، اما مسئله ریاضی از توجه متراکم به توجه پراکنده تغییر یافته است.

این طراحی مشترک در سطح سیستم ثابت می‌کند که سرعت واقعی (Wall-clock speed) تنها توسط تعداد FLOPها تعیین نمی‌شود. با هم‌سو کردن زمان‌بندی ریاضی با سلسله‌مراتب حافظه سخت‌افزار، FlashAttention پنجره‌های متنی طولانی را از نظر محاسباتی کاربردی می‌کند.

گام بعدی شما

  • اگر از PyTorch استفاده می‌کنید، از torch.nn.functional.scaled_dot_product_attention بهره ببرید تا به‌طور خودکار از هسته‌های FlashAttention استفاده شود.
  • در زمان استنتاج (Inference)، مقدار dropout را حتماً روی ۰.۰ تنظیم کنید تا از محاسبات اضافی جلوگیری شود.
  • برای مدل‌هایی با پنجره متنی بالای ۳۲ هزار توکن، بررسی کنید که آیا سخت‌افزار شما از FlashAttention-2 یا ۳ پشتیبانی می‌کند تا از گلوگاه HBM رها شوید.

اما داستان سخت‌افزاری این تحول حتی شگفت‌انگیزتر است — به تحلیل ما درباره‌ی تراشه‌های Blackwell مراجعه کنید. در همین راستا، برای مشاهده اینکه چگونه سخت‌افزارهای تخصصی مانند سخت‌افزار Cerebras سرعت استنتاج GPT-5.6 Sol را به ۷۵۰ توکن بر ثانیه رسانده‌اند، می‌توانید گزارش ما را بخوانید تا تفاوت رویکردهای معماری در حذف گلوگاه‌ها را درک کنید.

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

این فناوری با حذف گلوگاه حافظه، امکان پردازش متون بسیار طولانی را بدون نیاز به سخت‌افزارهای نجومی فراهم می‌کند. اعتبار این روش از پذیرش آن در هسته‌ی اصلی PyTorch و استفاده در تمام مدل‌های پیشرو (SOTA) تأیید شده است.

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

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

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

جایگزینی پیچیدگی حافظه با محاسبات اضافی در FlashAttention، یک چرخش پارادایمی در بهینه‌سازی مدل‌هاست. این رویکرد ثابت می‌کند که در عصر تراشه‌های مدرن، «هزینه جابه‌جایی داده» بسیار گران‌تر از «هزینه محاسبه» است. در نتیجه، آینده‌ی بهینه‌سازی مدل‌های زبانی نه در کاهش پارامترها، بلکه در طراحی الگوریتم‌های IO-aware نهفته است.

منابع

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

گفتگو

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

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

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

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

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

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

دات‌هوش

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

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