FlashAttention-2 (TH)

From Systems analysis Wiki
Jump to navigation Jump to search

FlashAttention-2 — คืออัลกอริทึมที่ได้รับการปรับปรุงซึ่งออกแบบมาเพื่อคำนวณกลไก attention ในโมเดลภาษาขนาดใหญ่ (LLM) อัลกอริทึมนี้พัฒนาโดย Tri Dao และนักวิจัยจากมหาวิทยาลัยสแตนฟอร์ด และได้รับการนำเสนอในเดือนกรกฎาคม ปี 2023[1] เป้าหมายหลักของมันคือการเร่งความเร็วการฝึกและการอนุมาน (inference) ของโมเดล transformer อย่างมีนัยสำคัญ ด้วยการใช้ทรัพยากรฮาร์ดแวร์ GPU อย่างมีประสิทธิภาพมากขึ้น ในขณะที่ยังคงรักษาความเหมือนกันอย่างสมบูรณ์กับการคำนวณของกลไก attention มาตรฐาน กล่าวคือ ไม่สูญเสียความแม่นยำ

FlashAttention-2 เป็นการพัฒนาต่อยอดจากอัลกอริทึม FlashAttention ที่นำเสนอโดยทีมเดียวกันในปี 2022 เวอร์ชันใหม่นี้แก้ปัญหาการโหลด GPU ที่ไม่สมบูรณ์ซึ่งพบในรุ่นก่อนหน้า และบรรลุความเร็วที่เพิ่มขึ้นเกือบสองเท่าเมื่อเทียบกับเวอร์ชันแรก

บริบทและที่มา: ปัญหา attention ใน transformer

กลไก self-attention มาตรฐานเป็นจุดคอขวดเมื่อทำงานกับลำดับข้อความที่ยาวใน transformer ความซับซ้อนในการคำนวณและการใช้หน่วยความจำเพิ่มขึ้นแบบ กำลังสอง (O(N²)) ตามความยาวของลำดับ (N) ซึ่งกำหนดข้อจำกัดร้ายแรงต่อความยาว context สูงสุดและความสามารถในการขยายขนาดของ LLM[1]

เพื่อแก้ปัญหานี้ อัลกอริทึม FlashAttention จึงถูกนำเสนอในปี 2022[2] แนวคิดหลักของมันได้แก่:

  • การคำนึงถึงลำดับชั้นหน่วยความจำของ GPU (IO-awareness): อัลกอริทึมลดการดำเนินการอ่าน/เขียนที่มีต้นทุนสูงระหว่างหน่วยความจำ GPU ที่ช้า (HBM) และหน่วยความจำสถิตที่เร็ว (SRAM) บนชิป
  • การประมวลผลแบบบล็อก (tiling): การคำนวณถูกแบ่งออกเป็นบล็อกเล็กๆ (tile) ซึ่งประมวลผลใน SRAM ที่เร็ว ทำให้หลีกเลี่ยงการสร้างเมทริกซ์ attention เต็มรูปแบบในหน่วยความจำได้

สิ่งนี้ทำให้บรรลุการเติบโตของการใช้หน่วยความจำแบบ เชิงเส้น (O(N)) และความเร็วที่เพิ่มขึ้น 2–4 เท่าเมื่อเทียบกับการนำไปใช้งานมาตรฐาน[2] FlashAttention ได้รับการยอมรับอย่างแพร่หลายและช่วยให้เกิดโมเดลที่มี context ที่ยาวขึ้นอย่างมีนัยสำคัญ เช่น จาก 2–4 พัน token (GPT-3) เป็น 128,000 token (GPT-4) และมากกว่านั้น[3] ตัวอย่างเช่น ในโมเดล Falcon-40B การใช้ FlashAttention ทำให้ inference เร็วขึ้น 3 เท่า และประสิทธิภาพการสร้างโดยรวมเร็วขึ้น 5 เท่าเมื่อเทียบกับ GPT-3[4]

การพัฒนาและเป้าหมายของ FlashAttention-2

แม้จะประสบความสำเร็จ FlashAttention เวอร์ชันแรกยังไม่ได้ใช้ทรัพยากรการคำนวณของ GPU อย่างเต็มที่ บนการ์ดจอ NVIDIA A100 ประสิทธิภาพอยู่ที่เพียง 25–40% ของค่าสูงสุดทางทฤษฎี (FLOPs/s)[1] สาเหตุหลักคือการโหลด Streaming Multiprocessors ที่ไม่เหมาะสมและการดำเนินการกับ shared memory ที่มีเกินความจำเป็น[5]

เป้าหมายของ FlashAttention-2 คือการเร่งความเร็วการคำนวณให้มากขึ้นโดยการกระจายงานแบบขนานที่มีประสิทธิภาพมากขึ้นและลดการดำเนินการเสริม อัลกอริทึมได้รับการเขียนใหม่ทั้งหมดโดยใช้ primitive ระดับต่ำของไลบรารี NVIDIA CUTLASS 3.x เพื่อให้ได้ประสิทธิภาพสูงสุด[6]

สถาปัตยกรรมทางเทคนิคและหลักการทำงาน

FlashAttention-2 นำเสนอการปรับปรุงสำคัญสามประการเพื่อเพิ่ม parallelism และประสิทธิภาพ[1]:

1. การลดการดำเนินการที่ไม่ใช่เมทริกซ์

อัลกอริทึมลดจำนวนการดำเนินการเลขทศนิยมเสริมที่ไม่ใช่การคูณเมทริกซ์ (non-matmul FLOPs) เนื่องจาก tensor core ของ GPU ได้รับการปรับให้เหมาะสมสำหรับการดำเนินการเมทริกซ์ (GEMM) โดยเฉพาะและรันได้เร็วกว่าถึง 16 เท่า การเปลี่ยนแปลงนี้จึงทำให้สามารถใช้บล็อก GPU ที่มีประสิทธิภาพสูงสุดในช่วงเวลาส่วนใหญ่ได้

2. Parallelism ที่ได้รับการปรับปรุง

ใน FlashAttention ดั้งเดิม การทำงานบน "head" attention หนึ่งตัวไม่ได้รับการกระจายแบบขนาน ซึ่งนำไปสู่การว่างงานเมื่อลำดับยาวและ batch size เล็ก FlashAttention-2 นำเสนอ inter-block parallelism: บัดนี้การคำนวณสำหรับ attention head หนึ่งตัวถูกกระจายระหว่าง Streaming Multiprocessors ต่างๆ ของ GPU ซึ่งเพิ่มการโหลดอย่างมีนัยสำคัญ

3. การแบ่งงานที่ได้รับการปรับปรุงภายในบล็อก

ในระดับของบล็อกการคำนวณหนึ่งบล็อก งานถูกกระจายใหม่ระหว่างกลุ่มของ thread (warp) เพื่อลดการแลกเปลี่ยนข้อมูลผ่าน shared memory สิ่งนี้ลดจำนวนการดำเนินการอ่าน/เขียนที่ซ้ำซ้อนซึ่งจำเป็นสำหรับการ normalize Softmax

ประสิทธิภาพและความมีประสิทธิผล

ด้วยการปรับปรุงทางสถาปัตยกรรม FlashAttention-2 แสดงให้เห็นการเพิ่มประสิทธิภาพอย่างมีนัยสำคัญ:

  • ความเร็วเพิ่มขึ้นสองเท่า: อัลกอริทึมทำงานเร็วกว่าประมาณ 2 เท่า เมื่อเทียบกับ FlashAttention เวอร์ชันแรก[1]
  • การใช้งาน GPU สูง: บน GPU NVIDIA A100 บรรลุ 50–73% ของ throughput สูงสุดทางทฤษฎี (TFLOPs) ซึ่งใกล้เคียงกับประสิทธิภาพของการดำเนินการคูณเมทริกซ์ที่ได้รับการปรับปรุง (GEMM)[1]
  • ความเร็วในการคำนวณที่สร้างสถิติใหม่:
    • บน GPU A100 บรรลุความเร็วสูงถึง 225 TFLOP/s ในรอบการฝึกแบบ end-to-end ของโมเดลประเภท GPT ซึ่งสอดคล้องกับ 72% การใช้งาน compute unit เพื่อเปรียบเทียบ attention มาตรฐานในเงื่อนไขเดียวกันโหลด GPU ต่ำกว่า 100 TFLOP/s[7]
    • บน GPU H100 ประสิทธิภาพสูงถึง 335 TFLOP/s[7]

การเพิ่มประสิทธิภาพดังกล่าวทำให้สามารถ เช่น ฝึกโมเดลที่มี context window 16,000 token ในเวลาเดียวกับที่เคยต้องใช้สำหรับ window 8,000 token[5] สิ่งสำคัญคืออัลกอริทึมยังคง แม่นยำ และเป็น deterministic ดังนั้นการใช้งานมันจึงไม่ส่งผลต่อคุณภาพการทำนายของโมเดล[8]

การประยุกต์ใช้และการผสานรวมในระบบนิเวศ

FlashAttention-2 กลายเป็นเครื่องมือมาตรฐานในระบบนิเวศ LLM อย่างรวดเร็ว มันถูกผสานรวมเข้ากับ framework และไลบรารียอดนิยมหลายแห่ง:

  • PyTorch: รองรับแบบ native
  • Hugging Face Transformers: การรองรับเปิดใช้งานด้วยพารามิเตอร์ `attn_implementation="flash_attention_2"` เมื่อโหลดโมเดล[9] เข้ากันได้กับสถาปัตยกรรมหลายสิบแบบ (GPT, Llama, Falcon, BERT และอื่นๆ)[10]
  • TensorRT-LLM, xFormers และ Triton: อัลกอริทึมได้รับการนำไปใช้งานสำหรับแพลตฟอร์มเหล่านี้ ซึ่งทำให้มีการประยุกต์ใช้อย่างกว้างขวาง[7]

การผสานรวมนี้ทำให้สามารถนำ FlashAttention-2 มาใช้ร่วมกับวิธีการปรับปรุงอื่นๆ ได้อย่างง่ายดาย เช่น การ quantization (GPTQ, QLoRA) และการ fine-tuning อย่างมีประสิทธิภาพ (PEFT)[9]

การเปรียบเทียบกับเวอร์ชันถัดไป

FlashAttention-3

การวิจัยในด้านการปรับปรุง attention ยังคงดำเนินต่อไป ในเดือนกรกฎาคม 2024 Tri Dao นำเสนอ FlashAttention-3 ซึ่งมุ่งเป้าไปที่การใช้ประโยชน์จากความสามารถของสถาปัตยกรรม GPU NVIDIA Hopper (H100/H200) นวัตกรรมสำคัญ[3]:

  • การรองรับ FP8: ใช้การคำนวณทศนิยม 8 บิตเพื่อเร่งความเร็วเพิ่มเติม
  • การดำเนินการแบบ asynchronous: ใช้ความสามารถแบบ asynchronous ของ GPU อย่างมีประสิทธิภาพมากขึ้น

FlashAttention-3 ให้ความเร็วเพิ่มขึ้น 1.5–2 เท่า เมื่อเทียบกับ FlashAttention-2 บน GPU H100 โดยบรรลุประสิทธิภาพสูงถึง 740 TFLOP/s (75% ของค่าสูงสุดทางทฤษฎี)[11]

เอกสารอ้างอิง

  • Dao, T. (2023). FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning. arXiv:2307.08691.
  • Dao, T. et al. (2022). FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness. arXiv:2205.14135.
  • 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.
  • Ye, Z. et al. (2025). FlashInfer: Efficient and Customizable Attention Engine for LLM Inference Serving. arXiv:2501.01005.
  • Chen, Y. et al. (2023). FlashDecoding++: Faster Large Language Model Inference on GPUs. arXiv:2311.01282.
  • Liu, Y. et al. (2024). FastAttention: Extending FlashAttention-2 to NPUs and Low-Resource GPUs. OpenReview: 76NYyOrnfk.
  • 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. (2024). FlashMask: Efficient and Rich Mask Extension of FlashAttention. OpenReview: rog0J435OO.
  • Abbott, V.; Zardini, G. (2025). FlashAttention on a Napkin: A Diagrammatic Approach to Deep Learning IO-Awareness. arXiv:2412.03317.

หมายเหตุ

  1. 1.0 1.1 1.2 1.3 1.4 1.5 Дао, Три. «FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning». arXiv:2307.08691 [cs.LG], 17 июля 2023 г. [1]
  2. 2.0 2.1 «Optimizing LLMs for Speed and Memory». Hugging Face Documentation. [2]
  3. 3.0 3.1 Дао, Три. «FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision». Tri Dao's Blog. [3]
  4. «FlashAttention vs FlashAttention-2 - an Analysis». E2E Networks Blog. [4]
  5. 5.0 5.1 «FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning». OpenReview. [5]
  6. «FlashAttention-2». Hazy Research, Stanford University. [6]
  7. 7.0 7.1 7.2 Дао, Три. «FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning» (PDF). arXiv:2307.08691. [7]
  8. Рашка, Себастьян. «Llama 2 and FlashAttention 2». Ahead of AI Magazine. [8]
  9. 9.0 9.1 Белькада, Юнес. «Faster and more memory efficient models with Flash Attention 2!». LinkedIn. [9]
  10. «GPU inference». Hugging Face Documentation. [10]
  11. Дао, Три, и др. «FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision». arXiv:2407.08608 [cs.LG], 11 июля 2024 г. [11]