FlashAttention (BN)

From Systems analysis Wiki
Jump to navigation Jump to search

FlashAttention — এটি একটি যুগান্তকারী অ্যালগরিদম যা attention মেকানিজম গণনার জন্য তৈরি করা হয়েছে। এটি বৃহৎ ভাষা মডেলের (LLM) প্রশিক্ষণ এবং inference উল্লেখযোগ্যভাবে ত্বরান্বিত করে, সম্পূর্ণ গণনার নির্ভুলতা বজায় রেখে। অ্যালগরিদমটি ২০২২ সালে ত্রি দাও (Tri Dao)-এর নেতৃত্বে স্ট্যানফোর্ড বিশ্ববিদ্যালয়ের গবেষক দল প্রথম উপস্থাপন করেন[1]

FlashAttention-এর মূল ধারণা হলো GPU-এর মেমোরি হায়ারার্কি বিবেচনায় রেখে গণনাকে পুনর্বিন্যস্ত করা, যা ধীর মেমোরিতে অ্যাক্সেসের সংখ্যা কমিয়ে আনে এবং স্ট্যান্ডার্ড attention মেকানিজমের প্রধান বাধা দূর করে।

স্ট্যান্ডার্ড attention-এর সমস্যা

Transformer-এ স্ট্যান্ডার্ড self-attention মেকানিজম নিম্নোক্ত সূত্র অনুযায়ী গণনা করা হয়: Attention(Q,K,V)=softmax(QKTdk)V যেখানে Q, K, V হলো যথাক্রমে query, key এবং value ম্যাট্রিক্স।

এই পদ্ধতির প্রধান সমস্যা হলো ক্রম দৈর্ঘ্য N-এর সাপেক্ষে বর্গীয় জটিলতা (O(N²)) — সময় এবং মেমোরি উভয় ক্ষেত্রে[1]। সাধারণ বাস্তবায়নে GPU মেমোরিতে N×N আকারের সম্পূর্ণ attention ম্যাট্রিক্স S গণনা করে সংরক্ষণ করতে হয়, যা দুটি গুরুত্বপূর্ণ সমস্যার সৃষ্টি করে:

  1. অধিক মেমোরি ব্যবহার: দীর্ঘ context নিয়ে কাজ করার সময় N×N ম্যাট্রিক্স সংরক্ষণ করা অসম্ভব হয়ে পড়ে।
  2. ইনপুট-আউটপুট (IO) অপারেশন: প্রধান বাধা হলো গাণিতিক অপারেশনের সংখ্যা নয়, বরং GPU-এর ধীর মেমোরিতে ক্রমাগত অ্যাক্সেস।

GPU মেমোরি হায়ারার্কি

সমস্যাটি বোঝার জন্য GPU-তে দুই ধরনের মেমোরির পার্থক্য জানা গুরুত্বপূর্ণ (NVIDIA A100 উদাহরণ হিসেবে):

  • SRAM (স্ট্যাটিক মেমোরি): ক্ষুদ্র পরিসরের (~২০ MB) দ্রুত অন-চিপ মেমোরি, বিশাল ব্যান্ডউইথ সহ (১৯ TB/s পর্যন্ত)।
  • HBM (উচ্চ-ব্যান্ডউইথ মেমোরি): বৃহৎ পরিসরের (৪০–৮০ GB) ধীর মেমোরি, তুলনামূলকভাবে অনেক কম ব্যান্ডউইথ সহ (প্রায় ১.৫ TB/s)[2]

এই অসামঞ্জস্যতার কারণে স্ট্যান্ডার্ড attention অ্যালগরিদম মেমোরি ব্যান্ডউইথ-সীমাবদ্ধ (memory-bound) হয়ে পড়ে, কারণ এটি ক্রমাগত ধীর HBM থেকে বড় ম্যাট্রিক্স পড়ে এবং লেখে, যা বিলম্বের প্রধান উৎস।

FlashAttention-এর মূল উদ্ভাবন

FlashAttention একটি IO-সচেতন (IO-aware) অ্যালগরিদম, যা HBM অ্যাক্সেস কমিয়ে সমস্যার সমাধান করে। এটি তিনটি মূল কৌশলের মাধ্যমে অর্জিত হয়।

টাইলিং এবং ব্লক-ভিত্তিক প্রক্রিয়াকরণ

সম্পূর্ণ ম্যাট্রিক্স একবারে প্রক্রিয়া করার পরিবর্তে, FlashAttention ইনপুট ম্যাট্রিক্স Q, K, V-কে ছোট ছোট ব্লকে (টাইল) ভাগ করে, যা দ্রুত SRAM-এ রাখা যায়। অ্যালগরিদম ক্রমানুসারে এই ব্লকগুলো লোড করে, সেগুলোর জন্য সমস্ত attention গণনা সম্পন্ন করে এবং চূড়ান্ত ফলাফল আপডেট করে — ধীর HBM-এ সম্পূর্ণ attention ম্যাট্রিক্স সংরক্ষণ না করেই[1]

অনলাইন Softmax গণনা

একটি মূল প্রযুক্তিগত অগ্রগতি হলো Softmax-এর "অনলাইন" গণনা। স্ট্যান্ডার্ড Softmax নরমালাইজেশনের জন্য ইনপুট ভেক্টরের সমস্ত উপাদান জানতে হয়। FlashAttention একটি পরিবর্তিত অ্যালগরিদম ব্যবহার করে যা আংশিকভাবে Softmax গণনা করতে পারে। এটি দুটি মধ্যবর্তী মান (বর্তমান সর্বোচ্চ এবং এক্সপোনেন্টের যোগফল) রক্ষণাবেক্ষণ করে, যা নতুন ব্লক প্রক্রিয়া করার সাথে সাথে আপডেট হয় — ফলে সম্পূর্ণ ম্যাট্রিক্সে একসাথে অ্যাক্সেস না করেও সঠিক ফলাফল পাওয়া যায়[2]

একটি CUDA কার্নেলে অপারেশন একত্রীকরণ

সমস্ত attention অপারেশন (ম্যাট্রিক্স গুণন QKᵀ, মাস্কিং, 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 কার্য (১K-৪K): ২.৪ গুণ ত্বরান্বিত[1]
  • মেমোরি সাশ্রয়: সঠিক বেসলাইন বাস্তবায়নের তুলনায় ২০ গুণ পর্যন্ত মেমোরি সাশ্রয়।
  • মডেলের মান উন্নতি: দীর্ঘ context নিয়ে কাজ করার সুযোগের কারণে, FlashAttention শুধু মান অক্ষুণ্ণ রাখে না, বরং উন্নত করে। উদাহরণস্বরূপ, GPT-2-এর perplexity ০.৭ পয়েন্ট উন্নত হয়েছে এবং দীর্ঘ নথি শ্রেণীবিভাগ কার্যে নির্ভুলতা ৬.৪ পয়েন্ট বৃদ্ধি পেয়েছে[1]

বিবর্তন এবং পরবর্তী উন্নয়ন

FlashAttention-এর সাফল্য একটি সম্পূর্ণ হার্ডওয়্যার-সচেতন অ্যালগরিদমের ধারার সূচনা করেছে।

FlashAttention-2 (২০২৩)

দ্বিতীয় সংস্করণ GPU সম্পদের আরও সম্পূর্ণ ব্যবহারের লক্ষ্যে তৈরি হয়েছিল। মূল FlashAttention-এ NVIDIA A100-এ দক্ষতা সর্বোচ্চের মাত্র ২৫–৪০% ছিল। FlashAttention-2 গণনা সমান্তরালতায় উন্নতি এনেছে, যা নিম্নোক্ত সুবিধা দিয়েছে[4]:

  • প্রথম সংস্করণের তুলনায় দ্বিগুণ গতি অর্জন।
  • GPU ব্যবহার তাত্ত্বিক সর্বোচ্চের ৫০–৭৩%-এ বৃদ্ধি।
  • ২৫৬ আকারের attention head-এর জন্য সমর্থন এবং Multi-Query Attention (MQA) আর্কিটেকচারের জন্য সমর্থন সম্প্রসারণ।

FlashAttention-3 (২০২৪)

তৃতীয় সংস্করণ বিশেষভাবে NVIDIA Hopper (H100) GPU আর্কিটেকচারের জন্য অপ্টিমাইজ করা হয়েছে[5]। এটি নতুন হার্ডওয়্যার সুবিধা ব্যবহার করে, যেমন Tensor Cores-এর অ্যাসিঙ্ক্রোনাসিটি এবং FP8 সমর্থন, যা নিম্নোক্ত অর্জন সম্ভব করেছে:

  • FlashAttention-2-এর তুলনায় আরও ১.৫–২ গুণ ত্বরান্বিত।
  • FP16-এ ৭৪০ TFLOPS এবং FP8-এ প্রায় ১.২ PFLOPS পর্যন্ত কার্যক্ষমতা অর্জন।

বিশেষায়িত সমাধান

FlashAttention-এর ধারণাগুলো অন্যান্য প্রকল্পে বিকশিত হয়েছে:

  • FlashInfer (২০২৫): LLM inference কার্যের জন্য বিশেষভাবে অপ্টিমাইজড কাস্টমাইজযোগ্য attention ইঞ্জিন। এটি স্ট্রিমিং জেনারেশন মোডে KV-cache-এর সাথে দক্ষ কাজের উপর দৃষ্টি নিবদ্ধ করে[6]
  • FlashMLA (২০২৪): Context cache সংকোচন সহ attention বাস্তবায়ন (latent attention), যা ন্যূনতম তথ্য হ্রাসে অত্যন্ত দীর্ঘ ক্রমে মেমোরি সাশ্রয় করতে দেয়[7]

শিল্প ও ইকোসিস্টেমে প্রভাব

FlashAttention একটি মৌলিক অগ্রগতি হিসেবে দ্রুত LLM-এর দক্ষ প্রশিক্ষণ এবং inference-এর জন্য শিল্পের মানদণ্ডে পরিণত হয়েছে। এটি PyTorch এবং Hugging Face-এর মতো মূল লাইব্রেরিতে একীভূত করা হয়েছে এবং বেশিরভাগ বৃহৎ ভাষা মডেলে (LLaMA, MPT, Falcon, Claude প্রভৃতি) ব্যবহৃত হয়।

FlashAttention এবং এর পরবর্তী সংস্করণগুলো ভাষা মডেলের context window বৃদ্ধিতে নির্ণায়ক ভূমিকা পালন করেছে: GPT-3-এর ২–৪ হাজার token থেকে GPT-4-এর ১২৮ হাজার token এবং পরীক্ষামূলক মডেলে কয়েক মিলিয়ন token পর্যন্ত[8]। অ্যালগরিদমটি transformer স্কেলিংয়ের পথে অন্যতম প্রধান বাধা দূর করেছে এবং AI অ্যাপ্লিকেশনের জন্য — দীর্ঘ নথি বিশ্লেষণ থেকে শুরু করে মাল্টিমোডাল বোঝাপড়া পর্যন্ত — নতুন সম্ভাবনার দ্বার উন্মোচন করেছে।

তথ্যসূত্র

  • GitHub-এ FlashAttention-এর অফিসিয়াল রিপোজিটরি

সাহিত্য

  • 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. 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. 2.0 2.1 Дао, Три, и др. «FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness». OpenReview. [২]
  3. «We're Training AI Twice as Fast This Year as Last». IEEE Spectrum. [৩]
  4. Дао, Три. «FlashAttention-2». tridao.me. [৪]
  5. «FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision». PyTorch Blog. [৫]
  6. «[2501.01005] FlashInfer: Efficient and Customizable Attention Engine for LLM Inference Serving». arXiv. [৬]
  7. «GitHub - deepseek-ai/FlashMLA: FlashMLA: Efficient MLA decoding kernels». GitHub. [৭]
  8. «The Evolution of Flash Attention: Revolutionizing Transformer Efficiency». Medium. [৮]