FlashAttention-3 (TH)
FlashAttention-3 — คืออัลกอริทึมสำหรับการปรับปรุงประสิทธิภาพกลไก attention ในโครงข่ายประสาทเทียมแบบ transformer ซึ่งได้รับการพัฒนาเพื่อใช้ประโยชน์สูงสุดจากความสามารถของฮาร์ดแวร์ GPU สถาปัตยกรรม NVIDIA Hopper (H100)[1] อัลกอริทึมนี้ได้รับการนำเสนอในปี 2024 โดยกลุ่มนักวิจัยจากบริษัท Colfax Research, Meta, NVIDIA, Georgia Tech, มหาวิทยาลัย Princeton และ Together AI งานวิจัยนี้ได้รับการตอบรับในการประชุม NeurIPS 2024 และได้รับการคัดเลือกเป็น spotlight[2]
FlashAttention-3 เป็นการพัฒนาในรุ่นที่สามของตระกูลอัลกอริทึม ต่อจาก FlashAttention (2022) และ FlashAttention-2 (2023) เป้าหมายหลักคือการเร่งความเร็วการฝึกและการอนุมาน (inference) ของโมเดลภาษาขนาดใหญ่ (LLM) อย่างมีนัยสำคัญ ในขณะที่ยังคงรักษาความแม่นยำของการคำนวณไว้
บทนำและภูมิหลัง
ปัญหาของกลไก attention
องค์ประกอบสำคัญของ transformer คือกลไก self-attention อย่างไรก็ตาม ความซับซ้อนในการคำนวณและการใช้หน่วยความจำเพิ่มขึ้นแบบ กำลังสอง (O(n²)) ตามความยาวของลำดับข้อมูลขาเข้า (n)[1] สิ่งนี้ก่อให้เกิด "คอขวด" ที่ร้ายแรง เนื่องจาก GPU สมัยใหม่ได้รับการปรับแต่งสำหรับการคูณเมทริกซ์ที่รวดเร็ว แต่การคำนวณฟังก์ชันเลขยกกำลัง (เช่น ใน Softmax) นั้นช้ากว่าหลายเท่า นอกจากนี้ ในการติดตั้งแบบเรียบง่าย หน่วยความจำ GPU จะต้องเก็บเทนเซอร์ attention ขนาดใหญ่ ซึ่งจำกัดความสามารถในการปรับขนาดของโมเดล
FlashAttention และ FlashAttention-2
เพื่อแก้ปัญหานี้ ในปี 2022 จึงได้มีการนำเสนอ FlashAttention ซึ่งลดปริมาณการเข้าถึงหน่วยความจำทั่วไปที่ช้า (HBM) โดยใช้สองเทคนิคหลัก:
- การประมวลผลแบบบล็อก (tiling): การคำนวณถูกแบ่งออกเป็นบล็อก (tile) ที่ประมวลผลในหน่วยความจำ on-chip ที่รวดเร็ว (SRAM)
- การรวมการดำเนินการ: การดำเนินการทั้งหมด (การคูณเมทริกซ์, Softmax) ดำเนินการในเคอร์เนล GPU เดียวโดยไม่มีการบันทึกผลลัพธ์ระหว่างกลางลงในหน่วยความจำทั่วไป
สิ่งนี้ทำให้สามารถลดความซับซ้อนของหน่วยความจำจากกำลังสองเป็น เชิงเส้น และเร่งการคำนวณได้ 2–4 เท่า
ในปี 2023 ได้มีการนำเสนอเวอร์ชันที่ปรับปรุงแล้ว — FlashAttention-2 ซึ่งปรับปรุงการประมวลผลแบบขนาน บน GPU สถาปัตยกรรม NVIDIA Ampere (A100) ได้บรรลุ ~70% ของประสิทธิภาพสูงสุดทางทฤษฎี[3] อย่างไรก็ตาม บนสถาปัตยกรรมใหม่กว่าอย่าง NVIDIA Hopper (H100) ประสิทธิภาพกลับต่ำกว่ามาก — ประมาณ 35%[1] ซึ่งเกิดจากการที่อัลกอริทึมไม่ได้ใช้ประโยชน์จากความสามารถฮาร์ดแวร์ใหม่ของ Hopper และนั่นคือแรงผลักดันให้เกิดการสร้าง FlashAttention-3
ความสามารถฮาร์ดแวร์ใหม่ของ GPU Hopper (H100)
สถาปัตยกรรม NVIDIA Hopper นำเสนอฟังก์ชันใหม่หลายอย่างที่ FlashAttention-3 นำมาใช้เพื่อให้บรรลุประสิทธิภาพสูงสุด[4]:
- WGMMA (Warpgroup Matrix Multiply-Accumulate): ชุดคำสั่งใหม่สำหรับ tensor core ที่ดำเนินการคูณเมทริกซ์ด้วยประสิทธิภาพที่เกือบสองเท่าเมื่อเทียบกับสถาปัตยกรรม Ampere
- TMA (Tensor Memory Accelerator): โมดูลฮาร์ดแวร์ที่เร่งการถ่ายโอนข้อมูลระหว่างหน่วยความจำทั่วไป (HBM) และหน่วยความจำร่วม (shared memory) TMA ดำเนินการคำนวณที่อยู่โดยอัตโนมัติ ลดภาระของหน่วยประมวลผล
- รูปแบบ FP8: การรองรับฮาร์ดแวร์สำหรับรูปแบบข้อมูลทศนิยมแบบ 8 บิต ซึ่งเพิ่มประสิทธิภาพทางทฤษฎีเป็นสองเท่าเมื่อเทียบกับ FP16 แต่มีความเสี่ยงที่จะสูญเสียความแม่นยำเนื่องจากช่วงไดนามิกที่จำกัด
นวัตกรรมทางเทคนิคของ FlashAttention-3
อัลกอริทึมนำเสนอวิธีการปรับปรุงประสิทธิภาพหลักสามประการที่ออกแบบมาเฉพาะสำหรับสถาปัตยกรรม Hopper[4]:
1. การประมวลผลแบบอะซิงโครนัสและการเชี่ยวชาญเฉพาะด้านของ warp
FlashAttention-3 ใช้หลักการ warp-specialization โดยกลุ่มเธรด (warps) ต่างๆ บน GPU มีความเชี่ยวชาญเฉพาะสำหรับงานที่แตกต่างกัน:
- Producer warps: โหลดข้อมูลจากหน่วยความจำทั่วไปโดยใช้ TMA
- Consumer warps: ดำเนินการคูณเมทริกซ์บน tensor core
ด้วยความเป็นอะซิงโครนัสของฮาร์ดแวร์ Hopper การดำเนินการเหล่านี้จะ ทับซ้อนกันในเวลา ในขณะที่กลุ่ม warp หนึ่งดำเนินการคำนวณ อีกกลุ่มหนึ่งจะโหลดข้อมูลสำหรับบล็อกถัดไปแบบขนาน แนวทางแบบ pipeline นี้ที่จัดระเบียบตามหลักการ "ปิงปอง" (ping-pong scheduling) ช่วยให้ซ่อนความล่าช้าจากการดำเนินการที่ช้า (เช่น Softmax) และโหลดโมดูลฟังก์ชันทั้งหมดของ GPU ได้อย่างเต็มที่
2. การลดการดำเนินการกับหน่วยความจำให้น้อยที่สุด
อัลกอริทึมยังคงแนวคิด tiling จากเวอร์ชันก่อนหน้า แต่ใช้ TMA อย่างจริงจังสำหรับการโหลดบล็อกข้อมูลถัดไปแบบอะซิงโครนัส ควบคู่ กับการคำนวณปัจจุบัน การถ่ายโอนข้อมูลจาก HBM ที่ช้าไปยัง SRAM ที่เร็วนั้นดำเนินการ "ในเงา" ของการคำนวณหลัก ทำให้ GPU ไม่ต้องรอข้อมูลนาน
3. ความแม่นยำต่ำ (FP8) พร้อมการลดความผิดพลาดของการควอนไทซ์
การเปลี่ยนไปใช้ FP8 จะเพิ่มความเร็วเป็นสองเท่า แต่อาจนำไปสู่การสูญเสียความแม่นยำอย่างมีนัยสำคัญเนื่องจากการควอนไทซ์ เพื่อแก้ปัญหานี้ นักพัฒนาได้นำวิธีการ incoherent processing มาใช้[4] หลักการของวิธีนี้มีดังนี้:
- ก่อนการคำนวณ attention เวกเตอร์คุณลักษณะ (query Q และ key K) จะถูกคูณด้วย เมทริกซ์ออร์โธโกนัลสุ่ม (เช่น เมทริกซ์ Hadamard)
- การแปลงนี้จะ "กระจาย" ค่าที่มีขนาดผิดปกติ (outlier) ไปยังพิกัดทั้งหมด ทำให้การแจกแจงของค่าเหล่านั้นสม่ำเสมอขึ้น
- หลังจากนั้นจึงทำการควอนไทซ์เป็น FP8 ซึ่งจะเกิดขึ้นด้วยความผิดพลาดที่น้อยลง
- เนื่องจากการแปลงเป็นแบบออร์โธโกนัล จึงไม่บิดเบือนผลลัพธ์สุดท้ายของ attention (QKᵀ) เนื่องจากผลกระทบของเมทริกซ์จะถูกหักล้างเมื่อคูณกัน
เทคนิคนี้ช่วยลดความผิดพลาดในการคำนวณ attention ใน FP8 ลงประมาณ 2.6 เท่า เมื่อเทียบกับการใช้ FP8 มาตรฐานโดยไม่มีการแปลง[4]
ประสิทธิภาพและความสำคัญ
การนำเทคนิคเหล่านี้มาใช้ทำให้ FlashAttention-3 บรรลุความเหนือกว่าเวอร์ชันก่อนหน้าบน GPU H100 อย่างมีนัยสำคัญ:
- เร็วขึ้น 1.5–2 เท่า เมื่อเทียบกับ FlashAttention-2
- การใช้งาน GPU สูง: บรรลุ ~75–85% ของประสิทธิภาพสูงสุดทางทฤษฎีของ H100
- ปริมาณงาน:
- สูงถึง 740–840 TFLOPS สำหรับความแม่นยำครึ่งหนึ่ง (FP16/BF16)
- สูงถึง 1.2–1.3 PFLOPS (petaflops) เมื่อใช้ความแม่นยำ 8 บิต (FP8)[2]
ประสิทธิภาพสูงของ FlashAttention-3 ส่งผลโดยตรงต่อการพัฒนาและการนำ LLM ไปใช้:
- การลดเวลาการฝึก: การเร่งความเร็ว attention 75–100% ช่วยลดเวลาการฝึกโมเดลอย่างมีนัยสำคัญ ซึ่งอาจใช้เวลาเป็นสัปดาห์หรือเดือน
- การขยายหน้าต่าง context: โมเดลสามารถประมวลผลลำดับที่ยาวขึ้นได้อย่างมีประสิทธิภาพ (หลายแสน token) ซึ่งมีความสำคัญสำหรับการวิเคราะห์เอกสารขนาดใหญ่หรือโค้ด[1]
- การใช้ทรัพยากรอย่างมีประสิทธิภาพ: ช่วยให้บรรลุประสิทธิภาพเดิมด้วย GPU จำนวนน้อยลง หรือได้รับความเร็วที่สูงขึ้นด้วยฮาร์ดแวร์เดิม ซึ่งช่วยลดต้นทุนการใช้งานโมเดล
ความพร้อมใช้งานและการผนวกรวม
ผู้เขียนได้เผยแพร่โค้ดต้นฉบับของ FlashAttention-3 ภายใต้ลิขสิทธิ์แบบเปิดบน GitHub[4] คาดว่าจะมีการผนวกรวมเข้ากับ framework การเรียนรู้เชิงลึกชั้นนำ เช่น PyTorch และไลบรารี Hugging Face Transformers ซึ่งจะทำให้เทคโนโลยีนี้เข้าถึงได้สำหรับนักพัฒนาและนักวิจัยในวงกว้าง เวอร์ชันก่อนหน้าได้กลายเป็นมาตรฐานในอุตสาหกรรมไปแล้ว และ FlashAttention-3 มีแนวโน้มจะสืบสานแนวโน้มนี้ต่อไป
ลิงก์
- ที่เก็บโค้ดอย่างเป็นทางการของ FlashAttention บน GitHub
- บล็อกของ Together AI พร้อมประกาศ FlashAttention-3
บรรณานุกรม
- Shah, J. et al. (2024). FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision. arXiv:2407.08608.
- 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.
- 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. arXiv:2410.01359.
- Abbott, V.; Zardini, G. (2025). FlashAttention on a Napkin: A Diagrammatic Approach to Deep Learning IO-Awareness. arXiv:2412.03317.
หมายเหตุ
- ↑ 1.0 1.1 1.2 1.3 «FlashAttention-3 unleashes the power of H100 GPUs for LLMs». VentureBeat. [1]
- ↑ 2.0 2.1 Шах, Джей, и др. «FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision». OpenReview. [2]
- ↑ Шах, Джей, и др. «FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision». arXiv:2407.08608v2 [cs.LG], 15 июля 2024 г. [3]
- ↑ 4.0 4.1 4.2 4.3 4.4 Шах, Джей, и др. «FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision». Together AI Blog. [4]