FlashAttention (FA)
FlashAttention — یک الگوریتم انقلابی برای محاسبه مکانیزم attention است که برای تسریع قابلتوجه آموزش و inference مدلهای زبانی بزرگ (LLM) با حفظ دقت کامل محاسبات طراحی شده است. این الگوریتم برای اولین بار در سال ۲۰۲۲ توسط تیمی از پژوهشگران دانشگاه استنفورد به رهبری تری دائو (Tri Dao) معرفی شد[1].
ایده کلیدی FlashAttention در بازسازی محاسبات با در نظر گرفتن سلسلهمراتب حافظه GPU نهفته است، که این امر امکان به حداقل رساندن تعداد دسترسیها به حافظه کُند و رفع گلوگاه اصلی مکانیزم attention استاندارد را فراهم میسازد.
مشکلات attention استاندارد
مکانیزم self-attention استاندارد در transformerها بر اساس فرمول زیر محاسبه میشود: که در آن Q، K، V به ترتیب ماتریسهای query، key و value هستند.
مشکل اصلی این رویکرد، پیچیدگی درجه دوم از نظر زمان و حافظه (O(N²)) نسبت به طول دنباله N است[1]. در پیادهسازی ساده، لازم است ماتریس کامل attention به نام S با ابعاد N×N در حافظه GPU محاسبه و ذخیره شود که به دو مشکل بحرانی منجر میگردد:
- مصرف بالای حافظه: ذخیره ماتریس N×N هنگام کار با context های طولانی غیرممکن میشود.
- عملیات ورودی/خروجی (IO): گلوگاه اصلی نه تعداد عملیات حسابی، بلکه دسترسیهای مکرر به حافظه کُند GPU است.
سلسلهمراتب حافظه GPU
برای درک این مشکل، تمایز بین دو نوع حافظه در GPU (به عنوان مثال NVIDIA A100) اهمیت دارد:
- SRAM (حافظه استاتیک): حافظه سریع درونتراشهای با حجم کم (~۲۰ مگابایت) و پهنای باند بسیار بالا (تا ۱۹ ترابایت بر ثانیه).
- HBM (حافظه پهنای باند بالا): حافظه کُند با حجم زیاد (۴۰ تا ۸۰ گیگابایت) و پهنای باند به مراتب کمتر (حدود ۱.۵ ترابایت بر ثانیه)[2].
این عدم تقارن، الگوریتم attention استاندارد را محدود به پهنای باند حافظه (memory-bound) میکند، زیرا این الگوریتم به طور مداوم ماتریسهای بزرگ را از HBM کُند میخواند و مینویسد که منشأ اصلی تأخیرهاست.
نوآوریهای کلیدی FlashAttention
FlashAttention یک الگوریتم آگاه به IO (IO-aware) است که مشکل را از طریق به حداقل رساندن دسترسیها به HBM حل میکند. این امر با سه تکنیک اصلی محقق میشود.
Tiling و پردازش بلوکی
به جای پردازش کل ماتریس به صورت یکجا، FlashAttention ماتریسهای ورودی Q، K، V را به بلوکهای کوچک (tile) تقسیم میکند که در SRAM سریع جای میگیرند. الگوریتم این بلوکها را به ترتیب بارگذاری کرده، تمام محاسبات attention را روی آنها انجام میدهد و نتیجه نهایی را بهروزرسانی میکند، بدون اینکه ماتریس کامل attention را در HBM کُند ذخیره کند[1].
محاسبه آنلاین Softmax
پیشرفت فنی کلیدی، محاسبه «آنلاین» Softmax بود. Softmax استاندارد برای نرمالسازی نیاز به دانستن تمام عناصر بردار ورودی دارد. FlashAttention از الگوریتم اصلاحشدهای استفاده میکند که محاسبه Softmax را به صورت بخشی ممکن میسازد. این الگوریتم دو مقدار میانی (حداکثر جاری و مجموع توانهای نمایی) را نگه میدارد که با پردازش بلوکهای جدید بهروزرسانی میشوند و امکان دستیابی به نتیجه دقیق بدون دسترسی به کل ماتریس در یک لحظه را فراهم میکنند[2].
ادغام عملیات در یک هسته CUDA
تمام عملیات attention (ضرب ماتریسی QKᵀ، masking، Softmax، ضرب در V) در یک هسته CUDA ادغامشده واحد (fused kernel) ترکیب میشوند. این امر تعداد عملیات خواندن/نوشتن در HBM را به شکل چشمگیری کاهش میدهد: به جای گذرهای مکرر روی کل ماتریس، الگوریتم یک بلوک را یک بار در SRAM بارگذاری میکند، تمام محاسبات را انجام میدهد و تنها نتیجه نهایی را مینویسد.
کارایی نظری و عملی
پیچیدگی و بهینگی
FlashAttention مصرف حافظه را از O(N²) به O(N) کاهش میدهد که مقیاسپذیری خطی را تضمین میکند. ثابت شده است که IO-پیچیدگی الگوریتم برای محاسبه attention در سلسلهمراتب دو سطحی حافظه از نظر نظری بهینه است، به این معنی که اجرای attention دقیق بدون تغییر سختافزار سریعتر از این ممکن نیست[3].
نتایج تجربی
نسخه اول FlashAttention بهبودهای قابل توجهی نشان داد:
- تسریع:
- BERT-large (طول ۵۱۲): ۱۵٪ تسریع در آموزش.
- GPT-2 (طول ۱K): تسریع ۳ برابری.
- وظایف Long-Range Arena (1K تا 4K): تسریع ۲.۴ برابری[1].
- صرفهجویی در حافظه: تا ۲۰ برابر صرفهجویی در حافظه در مقایسه با پیادهسازیهای پایه دقیق.
- بهبود کیفیت مدلها: با توجه به امکان کار با context های طولانیتر، FlashAttention نه تنها کیفیت مدلها را کاهش نمیدهد، بلکه آن را بهبود میبخشد. به عنوان مثال، perplexity مدل GPT-2 به اندازه ۰.۷ واحد بهبود یافت و دقت در وظایف دستهبندی اسناد طولانی ۶.۴ واحد افزایش پیدا کرد[1].
تکامل و توسعههای بعدی
موفقیت FlashAttention آغازگر یک سلسله کامل از الگوریتمهای سختافزارمحور شد.
FlashAttention-2 (2023)
نسخه دوم به منظور استفاده کاملتر از منابع GPU طراحی شده بود. در FlashAttention اصلی، بهرهوری روی NVIDIA A100 تنها ۲۵ تا ۴۰ درصد از حداکثر بود. FlashAttention-2 بهبودهایی در موازیسازی محاسبات معرفی کرد که این امکان را فراهم آورد[4]:
- دستیابی به تسریع دو برابری نسبت به نسخه اول.
- افزایش بهرهوری GPU تا ۵۰ تا ۷۳٪ از حداکثر نظری.
- گسترش پشتیبانی تا head های attention با اندازه ۲۵۶، و همچنین برای معماریهای Multi-Query Attention (MQA).
FlashAttention-3 (2024)
نسخه سوم به طور خاص برای معماری GPU NVIDIA Hopper (H100) بهینهسازی شده است[5]. این نسخه از قابلیتهای سختافزاری جدید مانند ناهمزمانی Tensor Cores و پشتیبانی از FP8 استفاده میکند که این امکان را فراهم آورده است:
- دستیابی به تسریع ۱.۵ تا ۲ برابری بیشتر نسبت به FlashAttention-2.
- دستیابی به عملکرد تا ۷۴۰ TFLOPS روی FP16 و نزدیک به ۱.۲ PFLOPS روی FP8.
راهحلهای تخصصی
ایدههای FlashAttention در پروژههای دیگر توسعه یافتند:
- FlashInfer (2025): موتور attention قابل تنظیم که به طور خاص برای وظایف inference در LLM بهینهسازی شده است. این موتور بر کار کارآمد با KV-cache در حالت تولید جریانی تمرکز دارد[6].
- FlashMLA (2024): پیادهسازی attention با فشردهسازی cache متنی (latent attention) که امکان صرفهجویی در حافظه برای دنبالههای بسیار طولانی با حداقل از دست دادن اطلاعات را فراهم میکند[7].
تأثیر بر صنعت و اکوسیستم
FlashAttention به یک پیشرفت بنیادی تبدیل شد و به سرعت به استاندارد صنعت برای آموزش و inference کارآمد LLM بدل گشت. این الگوریتم در کتابخانههای کلیدی مانند PyTorch و Hugging Face ادغام شده و در اکثر مدلهای زبانی بزرگ (LLaMA، MPT، Falcon، Claude و غیره) مورد استفاده قرار میگیرد.
دقیقاً FlashAttention و نسخههای بعدی آن نقش تعیینکنندهای در افزایش پنجرههای context مدلهای زبانی ایفا کردند: از ۲ تا ۴ هزار token (GPT-3) تا ۱۲۸ هزار token (GPT-4) و حتی تا میلیونها token در مدلهای آزمایشی[8]. این الگوریتم یکی از موانع اصلی در مسیر مقیاسبندی transformerها را برطرف کرد و افقهای جدیدی را برای کاربردهای هوش مصنوعی، از تحلیل اسناد طولانی تا درک چندوجهی، گشود.
پیوندها
- مخزن رسمی FlashAttention در GitHub
منابع
- Dao, T. et al. (2022). FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness. arXiv:2205.14135.
- Dao, T. (2023). FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning. arXiv:2307.08691.
- Shah, J. et al. (2024). FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-Precision. arXiv:2407.08608.
- Kwon, W. et al. (2023). Efficient Memory Management for Large Language Model Serving with PagedAttention. arXiv:2309.06180.
- Hong, K. et al. (2023). FlashDecoding++: Faster Large Language Model Inference on GPUs. arXiv:2311.01282.
- Ye, Z. et al. (2025). FlashInfer: Efficient and Customizable Attention Engine for LLM Inference Serving. arXiv:2501.01005.
- Dege, P. et al. (2025). FlashMLA-ETAP: Efficient Transpose Attention Pipeline for Accelerating MLA Inference on NVIDIA H20 GPUs. arXiv:2506.01969.
- Wang, G. et al. (2025). FlashMask: Efficient and Rich Mask Extension of FlashAttention. OpenReview wUtXB43Chi.
- Dao, T. et al. (2022). FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness (OpenReview version). OpenReview H4DqfPSibmx.
- Gholami, A. et al. (2024). FlashAttention on a Napkin: A Diagrammatic Approach to Deep Learning IO-Awareness. OpenReview pF2ukh7HxA.
یادداشتها
- ↑ 1.0 1.1 1.2 1.3 1.4 Дао, Три, и др. «FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness». arXiv:2205.14135 [cs.LG], 28 мая 2022 г. [۱]
- ↑ 2.0 2.1 Дао, Три, и др. «FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness». OpenReview. [۲]
- ↑ «We're Training AI Twice as Fast This Year as Last». IEEE Spectrum. [۳]
- ↑ Дао, Три. «FlashAttention-2». tridao.me. [۴]
- ↑ «FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision». PyTorch Blog. [۵]
- ↑ «[2501.01005] FlashInfer: Efficient and Customizable Attention Engine for LLM Inference Serving». arXiv. [۶]
- ↑ «GitHub - deepseek-ai/FlashMLA: FlashMLA: Efficient MLA decoding kernels». GitHub. [۷]
- ↑ «The Evolution of Flash Attention: Revolutionizing Transformer Efficiency». Medium. [۸]