الأنظمة والتحسين2022متقدم10 دقيقة قراءة

FlashAttention: انتباه دقيق وسريع وموفِّر للذاكرة بوعي حركة البيانات

FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness

Dao, T. · Fu, D. Y. · Ermon, S. · Rudra, A. · Ré, C. — NeurIPS

المشكلة

وقت حساب الانتباه الذاتي واستهلاكه للذاكرة يتزايدان تربيعياً مع طول السلسلة: مصفوفة الانتباه بحجم N×N لا بدّ من تجسيدها كاملةً في ذاكرة HBM على المعالج الرسومي. الأساليب التقريبية تضحّي بالدقة مقابل السرعة، لكنها نادراً ما تحقق تسريعاً فعلياً على أرض الواقع، لأنها تتجاهل عنق الزجاجة الحقيقي — نقل البيانات بين ذاكرة HBM البطيئة وذاكرة SRAM السريعة على الشريحة، وليس عدد العمليات الحسابية.

الإسهام

خوارزمية انتباه دقيقة واعية بحركة البيانات، تستخدم لتحميل كتل صغيرة من Q وK وV إلى ذاكرة SRAM السريعة، وتحسب الانتباه كتلةً كتلة عبر تقنية ، دون أن تُجسِّد مصفوفة الانتباه N×N في الذاكرة البطيئة أبداً. تُدمج خمس عمليات منفصلة في نواة GPU واحدة عبر دمج النوى، وتحلّ محل التخزين: يعيد حساب الانتباه آنياً بدل تخزين المصفوفة الكاملة. النتيجة: تسريع فعلي بمقدار 2–4 أضعاف، وتوفير في الذاكرة يصل إلى 20 ضعفاً، مع نتيجة رياضية مطابقة تماماً للانتباه المعياري.

الأثر

أصبح FlashAttention المعيار الفعلي لحساب الانتباه في المحوِّلات الإنتاجية. كل نموذج لغوي كبير حديث تقريباً — GPT-4 وClaude وLLaMA وغيرها — يعتمد FlashAttention أو أحد خلفائه (FlashAttention-2 وFlashAttention-3). بفضله أصبح توسيع نوافذ السياق من 2 ألف إلى أكثر من 128 ألف رمز أمراً ممكناً. والأهم أن فلسفته الواعية بحركة البيانات غيّرت تفكير المجتمع البحثي جذرياً: من التركيز على تقليل العمليات الحسابية إلى التركيز على تحسين — تحوّل حقيقي في أبحاث أنظمة تعلّم الآلة.

تخيّل طاهياً يُحضِّر مأدبة كبيرة. الطريقة المعتادة تقول: انقل كل المكوّنات من المستودع إلى طاولة المطبخ، رتّبها جميعاً، اطبخها، ثم أعِد النتائج. المشكلة أن الطاولة صغيرة، فمعظم المكوّنات تبقى في المستودع — والطاهي يقضي وقتاً في التنقل ذهاباً وإياباً أكثر مما يقضي في الطبخ نفسه.

وصفة FlashAttention مختلفة: أحضِر صينية صغيرة من المكوّنات كل مرة، اطبخها فوراً على الطاولة، سجِّل ملاحظة سريعة عمّا أنجزت، ثم كرِّر. المأدبة الكاملة لا يلزم أن تكون حاضرة دفعة واحدة. الطاهي لا يغادر المطبخ تقريباً، والمأدبة تُقدَّم في نصف الوقت — بالأطباق نفسها بالضبط.

عنق الزجاجة الحقيقي: نطاق الذاكرة لا القدرة الحسابية

يحسب المحوِّل بالصيغة softmax(QK/dk)V\text{softmax}(QK^\top / \sqrt{d_k})\,V. لسلسلة من NN رمزاً تنشأ مصفوفة درجات بحجم N×NN \times N — أي N2N^2 قيمة يجب تخزينها ومعالجتها. عند N=4096N = 4096 يعني ذلك 16 مليون مُدخَل لكل رأس انتباه في كل طبقة.

المفاجأة هنا أن عنق الزجاجة ليس الحساب الرياضي نفسه. المعالجات الرسومية الحديثة قادرة على تنفيذ تريليونات العمليات العائمة في الثانية. ما لا تستطيعه هو نقل البيانات بالسرعة الكافية. معالج A100 مثلاً يحسب بقدرة 312 TFLOPS، لكن نطاق ذاكرته HBM لا يتجاوز 2 تيرابايت/ثانية. حين تكون نسبة الحساب إلى نقل البيانات منخفضة — كما في العمليات العنصرية مثل والأقنعة — تبقى أنوية المعالج عاطلةً تنتظر وصول البيانات. هذا ما يسمّيه باحثو الأنظمة .

الانتباه المعياري يكتب مصفوفة N×NN \times N كاملةً إلى ، ثم يقرأها لحساب softmax، ثم يكتب نتيجة softmax، ثم يقرأها مرةً أخرى لضرب المصفوفة النهائي. كل رحلة ذهاب وإياب إلى الذاكرة البطيئة هي وقت مهدور. الأساليب التقريبية (كالانتباه المتناثر والانتباه الخطي) حاولت حل المشكلة بتقليل عدد العمليات — لكن بما أن العمليات لم تكن هي عنق الزجاجة أصلاً، فقد فشلت في أغلب الأحيان في تحقيق تسريع فعلي.

افتح في المختبر
انقر على كل مستوى من مستويات الذاكرة لترى حجمه وسرعته. لاحظ الفجوة الهائلة — نحو 100 ضعف — بين سرعة SRAM وسرعة HBM.
تستيقظ التجربة عند وصولك…

ثلاث أفكار في نواة واحدة

يجمع FlashAttention ثلاث تقنيات كلاسيكية من عالم الأنظمة — التبليط و وإعادة الحساب — في نواة GPU واحدة تحسب دون أن تكتب مصفوفة N×NN \times N إلى الذاكرة البطيئة أبداً.

1. التبليط — الفكرة بسيطة: قسِّم مصفوفات Q وK وV إلى كتل صغيرة تتسع في . حمِّل كتلة واحدة من K وV، احسب الانتباه الجزئي لكل كتل Q مقابلها، ثم انتقل إلى الكتلة التالية. بهذه الطريقة لا تُجمَّع مصفوفة الانتباه الكاملة أبداً. الأمر أشبه بقراءة كتاب ضخم فصلاً فصلاً بدل طباعة كل صفحاته دفعة واحدة على مكتبك.

2. دمج النوى — الانتباه المعياري يشغّل خمس نوى GPU منفصلة: ضرب مصفوفات ← قناع ← softmax ← ← ضرب مصفوفات. كل نواة تقرأ من HBM وتكتب إليها. FlashAttention يدمج الخمس في نواة واحدة: البيانات تدخل SRAM مرة واحدة، تُنفَّذ العمليات الخمس هناك، ولا يخرج إلى HBM إلا الناتج النهائي.

3. إعادة الحساب — في التمرير الخلفي، بدل تخزين مصفوفة الانتباه N×NN \times N لحساب ، يعيد FlashAttention حسابها آنياً من Q وK وV المخزنة أصلاً. يستبدل بذلك قليلاً من العمليات الإضافية بتوفير هائل في الذاكرة — كمن يعيد حساب عملية بسيطة بدل أن يملأ دفتراً ضخماً بكل النتائج الوسيطة.

افتح في المختبر
شاهد كيف يمرّ FlashAttention على كتل Q وK/V واحدة تلو الأخرى. مصفوفة N×N الكاملة لا تُخزَّن أبداً — فقط نتائج الكتل الصغيرة تتراكم في SRAM.
تستيقظ التجربة عند وصولك…

حيلة softmax المتدفق: البصيرة المفتاحية

التبليط يعمل بطبيعته مع ضرب المصفوفات لأنه عملية تجميعية — يمكنك تقسيمها إلى كتل ودمج النتائج لاحقاً. لكن softmax ليست تجميعية مباشرة: حساب softmax(xi)=exi/jexj\text{softmax}(x_i) = e^{x_i}/\sum_j e^{x_j} يتطلب معرفة كل الدرجات xjx_j للحصول على المقام. فكيف نُبلِّط عملية تحتاج رؤية الصف بأكمله؟

الحل الذكي: احتفظ بـإحصاءات جارية — قيمة عظمى جارية mm ومجموع أُسّي جارٍ \ell. كلما وصلت كتلة جديدة من الدرجات، حدِّث mm و\ell، ثم أعِد قياس الناتج المتراكم بمعامل تصحيح. بعد معالجة كل الكتل تكون النتيجة مطابقة رياضياً لحساب softmax على الصف الكامل دفعة واحدة.

تُعرف هذه التقنية بـsoftmax المتدفق (Milakov & Gimelshein, 2018). تسمح بتقسيم حساب softmax إلى أي عدد من الكتل، تُعالَج كلٌّ منها باستقلال في SRAM، مع نقل حالة جارية صغيرة فقط بين الكتل. تخيّل أنك تحسب متوسط درجات طلابك ورقة ورقة بدل انتظار جمع كل الأوراق — تحمل فقط المجموع الجاري وعدد الأوراق المصحَّحة.

m(new)=max(m(old),m~),(new)=em(old)m(new)(old)+em~m(new)~m^{(new)} = \max(m^{(old)},\, \tilde{m}), \qquad \ell^{(new)} = e^{m^{(old)} - m^{(new)}} \ell^{(old)} + e^{\tilde{m} - m^{(new)}} \tilde{\ell}
إعادة قياس softmax المتدفق — تحديث الإحصاءات الجارية عند وصول كتلة جديدةm = القيمة العظمى الجارية للصف · ℓ = المجموع الجاري للأُسّيات · المقادير ذات المَدّة (~) = إحصاءات الكتلة الجديدة · المُعامِلات الأُسّية تُصحِّح تقادم القيمة العظمى السابقة
افتح في المختبر
تابع خطوات softmax المتدفق: شاهد كيف تُحدَّث القيمة العظمى والمجموع الجاريان مع كل كتلة جديدة، وتأكد أن النتيجة النهائية تطابق softmax الصف الكامل.
تستيقظ التجربة عند وصولك…

الانتباه المعياري مقابل FlashAttention: حكاية الذاكرة

الفرق الجوهري يكمن في مكان حدوث الحساب. الانتباه المعياري كل نتيجة وسيطة في HBM — مصفوفة الدرجات S=QKS = QK^\top، ونتيجة softmax التي نسمّيها PP، وغالباً قناع الإسقاط العشوائي. كلها مصفوفات بحجم N×NN \times N. في المقابل، يُبقي FlashAttention كل الوسيطات في SRAM ولا يكتب إلى HBM إلا الناتج النهائي OO (بحجم N×dN \times d لا N×NN \times N).

استهلاك الذاكرة ينخفض من O(N2)O(N^2) إلى O(N)O(N). عند طول سلسلة 2 ألف رمز يعني ذلك توفيراً بمقدار 10 أضعاف، وعند 4 آلاف يصل إلى 20 ضعفاً. هذا بالتحديد ما يسمح للنماذج الحديثة باستخدام تتجاوز 128 ألف رمز — أمر كان مستحيلاً فعلياً مع الانتباه المعياري على المعالجات المتاحة.

افتح في المختبر
اسحب منزلق طول السلسلة وقارن استهلاك الذاكرة: الانتباه المعياري (تربيعي) مقابل FlashAttention (خطّي).
تستيقظ التجربة عند وصولك…

تعقيد الدخل والخرج: لماذا تقليل الرحلات أهم من تقليل العمليات

الإسهام النظري الأبرز لـ FlashAttention هو تحليل الانتباه من منظور — أي عدّ عمليات القراءة والكتابة على HBM بدلاً من عدّ العمليات الحسابية.

الانتباه المعياري يحتاج Θ(Nd+N2)\Theta(Nd + N^2) عملية وصول إلى HBM. أما FlashAttention فيحتاج Θ(N2d2M1)\Theta(N^2 d^2 M^{-1}) حيث MM هو حجم SRAM. وبما أن MM كبير نسبةً إلى d2d^2 (القيمة النموذجية لـ dd هي 64–128 وحجم MM النموذجي 100 كيلوبايت فأكثر)، فإن عدد عمليات الوصول إلى HBM يقلّ بفارق كبير.

الأهم من ذلك أن المؤلفين أثبتوا حداً أدنى نظرياً: لأي خوارزمية انتباه دقيقة تستخدم العمليات المعيارية، فإن تعقيد الدخل والخرج الذي حققه FlashAttention هو الأمثل حتى عامل ثابت ضمن مدى واسع من أحجام SRAM. بمعنى آخر، لا يمكن لأي خوارزمية أن تتفوق عليه ضمن نموذج العتاد هذا.

HBM accessesstandard=Θ(Nd+N2)HBM accessesflash=Θ ⁣(N2d2M)\text{HBM accesses}_{\text{standard}} = \Theta(Nd + N^2) \qquad \text{HBM accesses}_{\text{flash}} = \Theta\!\left(\frac{N^2 d^2}{M}\right)
مقارنة تعقيد الدخل والخرج — الانتباه المعياري مقابل FlashAttentionN = طول السلسلة · d = بُعد رأس الانتباه · M = حجم SRAM · يتفوق Flash لأن M ≫ d² مما يُكبِّر المقام ويُقلِّص عدد عمليات الوصول

الخوارزمية خطوة بخطوة

يعمل كالتالي. قبل الدخول في التفاصيل، افهم الإيقاع العام: حلقة خارجية تمرّ على كتل K/V، وحلقة داخلية تمرّ على كتل Q — لكل زوج، احسب بلاطة صغيرة من الانتباه، حدِّث إحصاءات softmax الجارية، وراكِم النتيجة في الناتج.

  1. قسِّم Q إلى كتل بحجم BrB_r، وK وV إلى كتل بحجم BcB_c.
  2. هيِّئ الناتج O=0O = 0، والقيمة العظمى الجارية m=m = -\infty، والمجموع الجاري =0\ell = 0.
  3. لكل كتلة K/V رقم jj: حمِّل KjK_j وVjV_j إلى SRAM.
  4. — لكل كتلة Q رقم ii: حمِّل QiQ_i إلى SRAM، واحسب Sij=QiKjS_{ij} = Q_i K_j^\top.
  5. — — احسب القيمة العظمى المحلية m~\tilde{m} والمجموع المحلي ~\tilde{\ell} من SijS_{ij}.
  6. — — حدِّث mm و\ell الجاريتين بصيغ softmax المتدفق.
  7. — — أعِد قياس الناتج الجاري OiO_i وأضِف إسهام الكتلة الجديدة.
  8. اكتب الناتج النهائي OO واحفظ \ell وmm (يحتاجهما التمرير الخلفي) إلى HBM.
افتح في المختبر
تابع خوارزمية FlashAttention كتلةً كتلة. شاهد تحديث الإحصاءات الجارية وتراكم الناتج — لا تتشكّل مصفوفة N×N أبداً.
تستيقظ التجربة عند وصولك…

الفكرة نفسها في شفرة برمجية

التمرير الأمامي لـ FlashAttention (بايثون مبسّطة)python

مبسَّط لإظهار الفكرة — ليس التنفيذ الحقيقي.

import numpy as np

def flash_attention(Q, K, V, block_size=64):
    """FlashAttention: انتباه دقيق مُبلَّط دون تجسيد مصفوفة N×N."""
    N, d = Q.shape
    O = np.zeros_like(Q)           # مُراكِم الناتج
    ell = np.zeros((N, 1))          # المجموع الجاري للأُسّيات
    m = np.full((N, 1), -np.inf)    # القيمة العظمى الجارية للصف

    # الحلقة الخارجية: تدفق كتل K/V (كتحميل صوانٍ إلى المطبخ)
    for j in range(0, N, block_size):
        Kj = K[j:j+block_size]
        Vj = V[j:j+block_size]

        # الحلقة الداخلية: كل كتلة Q تعالج كتلة K/V هذه
        for i in range(0, N, block_size):
            Qi = Q[i:i+block_size]
            Sij = Qi @ Kj.T / np.sqrt(d)     # بلاطة صغيرة من الدرجات

            # softmax المتدفق: تحديث القيمة العظمى والمجموع الجاريين
            m_new = np.maximum(m[i:i+block_size], Sij.max(axis=-1, keepdims=True))
            P = np.exp(Sij - m_new)           # أُسّ آمن بالقيمة العظمى الجديدة
            correction = np.exp(m[i:i+block_size] - m_new)

            # إعادة قياس الناتج القديم وإضافة الإسهام الجديد
            O[i:i+block_size] = correction * O[i:i+block_size] + P @ Vj
            ell[i:i+block_size] = correction * ell[i:i+block_size] + P.sum(axis=-1, keepdims=True)
            m[i:i+block_size] = m_new

    return O / ell   # التسوية النهائية

# هذا كل شيء. مصفوفة N×N لا توجد أبداً.
# نوى CUDA الحقيقية تفعل هذا في ذاكرة SRAM على الشريحة، لا في الذاكرة الرئيسية.

FlashAttention المتناثر الكُتَلي: تخطّي ما لا يهم

يمتد FlashAttention بشكل طبيعي إلى : إذا كان قناع التناثر يشير إلى أن الكتلة (i,j)(i, j) صفرية، تخطّاها بالكامل — لا تحميل ولا حساب ولا كتابة. الأمر بديهي ضمن إطار التبليط لأن كل كتلة هي أصلاً وحدة عمل مستقلة.

يوفر FlashAttention المتناثر الكتلي تسريعاً إضافياً بمقدار 2–4 أضعاف فوق FlashAttention الكثيف، بما يتناسب مع نسبة التناثر. هذا ما مكّن بسلاسل تصل إلى 64 ألف رمز — وأنتج أول محوِّل يحقق أداءً أفضل من الصدفة على معيار Path-256 (طول السلسلة 64 ألف رمز، بدقة 63.1%).

افتح في المختبر
بدِّل أنماط التناثر وشاهد أي الكتل يتخطاها FlashAttention. كلما زاد التناثر قلّت الكتل المطلوب حسابها وقلّت رحلات الوصول إلى HBM.
تستيقظ التجربة عند وصولك…

النتائج: تدريب أسرع وسياقات أطول ونماذج أفضل

أثر FlashAttention يظهر في السرعة وفيما تفتحه السياقات الأطول من إمكانيات:

  • تسريع شامل بنسبة 15% على BERT-large (طول السلسلة 512) مقارنة بالرقم القياسي لسرعة التدريب في MLPerf 1.1.
  • تسريع 3 أضعاف في حساب الانتباه على GPT-2 (طول السلسلة 1K).
  • تسريع 2.4 ضعف على مهام المدى الطويل (أطوال سلاسل 1K–4K).
  • تحسّن 0.7 في الحيرة على GPT-2 عند استخدام سياق أطول (نمذجة أفضل بحجم النموذج ذاته).
  • ارتفاع 6.4 نقطة في مهام تصنيف المستندات الطويلة.
  • أول أداء يتجاوز الصدفة على Path-X (16 ألف رمز، دقة 61.4%) وPath-256 (64 ألف رمز، دقة 63.1%) — معايير لم يستطع أي محوِّل التعامل معها سابقاً.

الأثر وما جاء بعده

  1. 2018

    softmax المتدفق

    قدّم Milakov وGimelshein تقنية softmax المتدفق للحساب المستقر بتمريرة واحدة — الأساس الرياضي الذي بنى عليه FlashAttention.

  2. 2022

    FlashAttention (هذه الورقة)

    جمع Dao وزملاؤه بين التبليط ودمج النوى وإعادة الحساب لتحقيق انتباه دقيق بتسريع 2–4 أضعاف واستهلاك ذاكرة خطي. نُشرت في NeurIPS 2022.

  3. 2023

    FlashAttention-2

    توازٍ أفضل وتوزيع عمل محسّن يقرّبان FlashAttention من الحدود القصوى للعتاد، ليصل إلى 50–73% من القدرة الحسابية النظرية لمعالج A100.

  4. 2024

    FlashAttention-3

    يستفيد من ميزات معالجات Hopper (الذاكرة غير المتزامنة وFP8) لرفع الإنتاجية أكثر، مقترباً من الأداء الأقصى لمعالج H100.

  5. 2024

    نوافذ سياق تتجاوز 128 ألف رمز

    أُطلقت نماذج مثل Claude وGPT-4 Turbo وGemini 1.5 بنوافذ سياق تمتد من 128 ألف إلى مليون رمز، وكلها أصبحت ممكنة بفضل عائلة خوارزميات FlashAttention.

الدرس الذي يقدّمه FlashAttention يتجاوز الانتباه بكثير. أي عملية مقيَّدة بالذاكرة — والإسقاط العشوائي و — تستفيد من الفلسفة الواعية بحركة البيانات نفسها. هذه الورقة غيّرت طريقة تفكير مجتمع أنظمة تعلّم الآلة في الأمثَلة: ابدأ من هرمية الذاكرة، لا من عدد العمليات الحسابية.

المرجعDao, Fu, Ermon, Rudra, Ré. FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness. NeurIPS, 2022.

مصطلحات هذه الورقة