FlashAttention-3 (TH)

From Systems analysis Wiki
Jump to navigation Jump to search

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] หลักการของวิธีนี้มีดังนี้:

  1. ก่อนการคำนวณ attention เวกเตอร์คุณลักษณะ (query Q และ key K) จะถูกคูณด้วย เมทริกซ์ออร์โธโกนัลสุ่ม (เช่น เมทริกซ์ Hadamard)
  2. การแปลงนี้จะ "กระจาย" ค่าที่มีขนาดผิดปกติ (outlier) ไปยังพิกัดทั้งหมด ทำให้การแจกแจงของค่าเหล่านั้นสม่ำเสมอขึ้น
  3. หลังจากนั้นจึงทำการควอนไทซ์เป็น FP8 ซึ่งจะเกิดขึ้นด้วยความผิดพลาดที่น้อยลง
  4. เนื่องจากการแปลงเป็นแบบออร์โธโกนัล จึงไม่บิดเบือนผลลัพธ์สุดท้ายของ 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. 1.0 1.1 1.2 1.3 «FlashAttention-3 unleashes the power of H100 GPUs for LLMs». VentureBeat. [1]
  2. 2.0 2.1 Шах, Джей, и др. «FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision». OpenReview. [2]
  3. Шах, Джей, и др. «FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision». arXiv:2407.08608v2 [cs.LG], 15 июля 2024 г. [3]
  4. 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]