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

بک‌اندهای SDPA در PyTorch؛ تفاوت عملکردی FlashAttention در برابر cuDNN و Math

·۱۹ تیر ۱۴۰۵۱۷ دقیقه مطالعه۱ بازدید
راهنما
پروفایل‌سازی در PyTorch (بخش ۳): توجه، همه چیز پروفایل‌سازی است
پروفایل‌سازی در PyTorch (بخش ۳): توجه، همه چیز پروفایل‌سازی است
اشتراک‌گذاری
واقعاً چه چیز جدید است؟

افشای مکانیسم دقیق کندی بک‌اند Math در PyTorch و اثبات اینکه حتی پیاده‌سازی ساده در-جا (in-place) می‌تواند از یک بک‌اند رسمی اما غیربهینه، ۳.۷ برابر سریع‌تر باشد.

یک خط کد در PyTorch می‌تواند مدل شما را شتاب ببخشد یا عملکرد آن را تا ۴ برابر تخریب کند. این تضاد تکان‌دهنده در سری تحلیل‌های پروفایلینگ شرکت Hugging Face که در ۱۰ ژوئیه ۲۰۲۶ منتشر شد، آشکار شد؛ جایی که مشخص شد انتخاب بک‌اند برای scaled_dot_product_attention (SDPA) به‌طور کلی نحوه مدیریت مکانیزم توجه (Attention) توسط GPU را تغییر می‌دهد.

برای توسعه‌دهندگانی که مدل‌های زبانی بزرگ (LLM) — مثل کتابخانه‌داری که میلیاردها صفحه را خوانده و حالا با همان لحن جواب می‌دهد — یا مدل‌های انتشار را می‌سازند، توجه گلوگاه اصلی است چون پیچیدگی زمانی آن به‌صورت درجه‌دوم رشد می‌کند. اگرچه APIهای سطح بالا این پیچیدگی را می‌پوشانند، اما اجرای سخت‌افزاری در بک‌اندهای مختلف کاملاً متفاوت است. این موضوع صرفاً مربوط به کارایی نرم‌افزاری نیست، بلکه به نحوه جابه‌جایی داده‌ها بین حافظه پهنای‌باند بالا (HBM) و رجیسترهای داخلی تراشه بازمی‌گردد. طبق گزارش Hugging Face، این آزمایش‌ها روی یک GPU مدل NVIDIA A100-SXM4-80GB انجام شده است.

همان‌طور که در تحلیل‌های قبلی ما درباره بهینه‌سازی حافظه در مدل‌های ترنسفورمر اشاره کردیم، مدیریت ترافیک داده بین VRAM و هسته‌های پردازشی، تعیین‌کننده نهایی سرعت است.

پیاده‌سازی ساده و هزینه پنهان حافظه

ساخت مکانیزم توجه از قطعات اولیه (مانند ضرب ماتریسی، مقیاس‌بندی، ماسک‌گذاری و سافت‌مکس) یک نشت عملکردی جدی را فاش می‌کند. در یک پیاده‌سازی ساده، PyTorch ابتدا امتیازها را محاسبه کرده، آن‌ها را مقیاس می‌کند و سپس یک ماسک علی (Causal Mask) اعمال می‌کند.

بر اساس مستندات تحلیل شده، استفاده از دستور masked_fill باعث ایجاد یک هسته کپی حافظه (Memcpy) پنهان در GPU می‌شود. دلیل این اتفاق آن است که PyTorch اغلب یک کپی از تنسور می‌گیرد و عملیات را روی کپی انجام می‌دهد، نه روی نسخه اصلی.

پروفایلینگ در PyTorch (بخش ۳): توجه تنها چیزی است که پروفایل می‌کنید

با تغییر ساده به نسخه در-جا یا in-place (یعنی masked_fill_ که با یک خط تیره در پایان شناخته می‌شود)، هسته Memcpy کاملاً حذف می‌شود. ردپای پروفایلر نشان می‌دهد که نسخه in-place عملیات‌های CPU بسیار کمتری را در مرحله ماسک‌گذاری اجرا می‌کند.

پروفایلینگ در PyTorch (بخش ۳): توجه، همه چیز پروفایلینگ است

این تغییر کوچک در هر لایه تکرار می‌شود و در یک مدل ترنسفورمر مدرن با ده‌ها لایه، مقدار قابل‌توجیهی از زمان و حافظه را ذخیره می‌کند. این روش به شرطی ایمن است که برنامه تحت torch.no_grad() اجرا شود؛ زیرا در این حالت نیازی به ذخیره مقادیر اصلی برای پاس بازگشتی (Backward Pass) نیست و بازنویسی داده‌ها مشکلی ایجاد نمی‌کند.

بک‌اند Math: تله‌ای برای عملکرد

وقتی توسعه‌دهندگان از F.scaled_dot_product_attention با تثبیت روی بک‌اند 'math' استفاده می‌کنند، با نتیجه‌ای شوکه کننده روبرو می‌شوند. به‌رغم اینکه فقط یک خط کد است، بک‌اند math حدود ۳.۷ برابر کندتر از پیاده‌سازی ساده in-place عمل می‌کند. در جدول پروفایلر، میانگین زمان CUDA برای عملیات *_fwd از ۱.۹۵۵ میلی‌ثانیه به ۷.۲۳۹ میلی‌ثانیه جهش می‌کند.

پروفایلینگ در PyTorch (بخش ۳): توجه، همه چیز پروفایلینگ است

پروفایلر فاش می‌کند که بک‌اند math در هر پاس پیشرو ۲۰ هسته GPU را فراخوانی می‌کند، در حالی که نسخه ساده تنها ۵ تا ۶ هسته دارد. این سقوط عملکرد به سه دلیل رخ می‌دهد:

۱. خالی ماندن هسته‌های تنسور (Tensor Cores)

  • ارتقای دقت به FP32: بک‌اند math برای دقت عددی بیشتر، تنسورها را به FP32 ارتقا می‌دهد. این کار حجم داده‌های جابه‌جایی را حتی برای ورودی‌های bf16 دو برابر می‌کند.
  • بازگشت به هسته‌های CUDA: به دلیل تبدیل به FP32، مدل از هسته‌های تخصصی Tensor Core عبور کرده و به هسته‌های general-purpose قدیمی‌تر بازمی‌گردد. نام هسته‌ها در ردپا به‌جای امضاهای Tensor Core، عبارت sgemm را نشان می‌دهد.

۲. تجسم ماسک (Mask Materialization)

  • ساخت لحظه‌ای: پرچم is_causal=True کار ماسک‌گذاری را حذف نمی‌کند، بلکه فقط جابه‌جا می‌کند. بک‌اند math در هر فراخوانی، ماسک را از نو می‌سازد.
  • خط لوله CPU: در مسیر CPU، توالی دستورات aten::ones و aten::tril برای ساخت یک ماتریس مثلثی پایین دیده می‌شود که در نهایت به یک فشار پردازشی روی GPU منجر می‌گردد.

۳. سافت‌مکس ایمن (Safe Softmax)

  • جلوگیری از NaN: برای جلوگیری از تولید مقادیر نامعتبر (NaN) در ردیف‌های کاملاً ماسک‌شده، بک‌اند math از aten::_safe_softmax استفاده می‌کند که باعث اضافه شدن هسته‌های پردازشی اضافی در ردپا می‌شود.

هسته‌های ادغام‌شده: بک‌اندهای Efficient و Flash

برای فرار از پراکندگی ۲۰ هسته‌ای در بک‌اند math، PyTorch از بک‌اندهای بهینه مانند FLASH_ATTENTION و EFFICIENT_ATTENTION استفاده می‌کند. این‌ها تمام عملیات‌های اولیه را در یک تک‌هسته (Fused Kernel) ادغام می‌کنند.

بک‌اند 'efficient' (مبتنی بر xformers) از هسته fmha_cutlassF استفاده می‌کند که روی CUTLASS ساخته شده، در حالت bfloat16 اجرا می‌شود و داده‌ها را در رجیسترهای داخلی نگه می‌دارد تا سرعت در GPUهای نسل Ampere به حداکثر برسد.

توجه برق‌آسا (FlashAttention-2) یک گام فراتر می‌رود و ترافیک HBM را هدف قرار می‌دهد. هزینه اصلی توجه نه ضرب ماتریسی، بلکه ترافیک رفت‌وبرگشت مداوم به حافظه برای خواندن و نوشتن ماتریس امتیازات است. برای یک توالی ۴۰۹۶، این یعنی جابه‌جایی حدود ۱۶ میلیون عدد برای هر سر توجه.

Flash از ترفندی به نام «سافت‌مکس آنلاین» استفاده می‌کند تا روی K و V به‌صورت تکه‌های کوچک (Tiles) حرکت کند. در این حالت، ماتریس کامل امتیازات هرگز در HBM نوشته نمی‌شود و فقط روی تراشه می‌ماند.

ردپای «غلط» پروفایلر در Flash

به‌طور جالب، پروفایلر اشغال حافظه (Occupancy) پایینی (حدود ۱۳٪) را برای FlashAttention گزارش می‌کند. این به معنای بهینه‌سازی差 نیست. Flash به‌جای نگه داشتن تعداد زیادی Warp برای پنهان کردن تأخیر، بر بازاستفاده‌ی خام از داده‌ها تمرکز دارد.

  • سنگینی منابع: Flash مقدار عظیمی از حافظه مشترک (Shared Memory) و رجیسترهای هر رشته را اشغال می‌کند.
  • تعداد Warps: در یک SM نسل Ampere، به‌دلیل مصرف بالای رجیسترها، تنها ۲ بلوک جای می‌گیرند که منجر به عدد ۱۳٪ می‌شود.

رویکرد cuDNN: تولید لحظه‌ای هسته

برخلاف Flash، بک‌اند NVIDIA cuDNN هسته‌هایی را در زمان اجرا (Runtime) متناسب با ابعاد مسئله تولید می‌کند. نام‌های طولانی هسته‌ها در ردپا (مانند cudnn_generated_...) گواه این رویکرد است.

ویژگی‌های کلیدی cuDNN عبارتند از:

  • تنظیم Knob: بر اساس ابعاد تنسور، یکی از پیکربندی‌های پیش‌تنظیم (Knobs) را انتخاب می‌کند.
  • چیدمان مستقیم: داده‌ها را مستقیماً با ساختار [B, H, S, D] مصرف کرده و چهار عملیات aten::transpose را که در Flash دیده می‌شود، حذف می‌کند.
  • فراخوانی سطح درایور: از cuLaunchKernelEx استفاده می‌کند که باعث می‌شود CUPTI نتواند اشغال حافظه را به‌درستی اندازه‌گیری کند و گاهی عدد ۰٪ را گزارش دهد.

این انعطاف‌پذیری هزینه‌ای روی CPU دارد. طبق داده‌های پروفایلر، cuDNN حدود ۲۱۴ میکروثانیه را صرف انتخاب بهترین Knob می‌کند، در حالی که این زمان برای Flash حدود ۱۳۸ میکروثانیه است.

گام بعدی شما

  • اگر از مدل‌های سفارشی در PyTorch استفاده می‌کنید، حتماً خروجی torch.profiler را بررسی کنید تا مطمئن شوید بک‌اند 'math' به‌طور ناخواسته فعال نشده است.
  • برای کاهش تأخیر در محیط Production، اولویت خود را به ترتیب FlashAttention-2 و سپس cuDNN قرار دهید.
  • در پیاده‌سازی‌های دستی، همیشه از نسخه‌های in-place مانند masked_fill_ استفاده کنید تا کپی‌های پنهان حافظه حذف شوند.

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

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

این یافته‌ها بر اعتبار فنی توسعه‌دهندگان تأکید می‌کند که صرفاً به APIهای سطح بالا تکیه نکنند. اشتباه در انتخاب بک‌اند می‌تواند هزینه‌های استنتاج و زمان آموزش را به‌طور بی‌دلیل افزایش دهد.

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

برای توسعه‌دهندگانی که در ایران با محدودیت منابع GPU (مانند دسترسی به تعداد کم A100) سر و کار دارند، استفاده از FlashAttention-2 تنها راه کاهش هزینه‌های پردازشی و افزایش سرعت استنتاج در مدل‌های محلی است.

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

این تحلیل ثابت می‌کند که در دنیای مدل‌های زبانی، «سادگی کد» لزوماً به معنای «بهینگی اجرا» نیست. جابه‌جایی تمرکز از بهینه‌سازی محاسبات (FLOPs) به بهینه‌سازی جابه‌جایی داده‌ها (Memory Bound)، پارادایم جدید توسعه مدل‌های مقیاس‌بزرگ است. در واقع، FlashAttention با پذیرفتن اشغال حافظه کمتر (Occupancy)، کارایی واقعی را از طریق کاهش ترافیک HBM به دست آورده است.

منابع

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

گفتگو

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

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

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

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

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

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

دات‌هوش

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

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