FlashAttention (TH)

From Systems analysis Wiki
Jump to navigation Jump to search

FlashAttention — คืออัลกอริทึมปฏิวัติวงการสำหรับการคำนวณกลไก attention ที่พัฒนาขึ้นเพื่อเร่งความเร็วในการฝึกและการ inference ของโมเดลภาษาขนาดใหญ่ (LLM) อย่างมีนัยสำคัญ พร้อมรักษาความแม่นยำในการคำนวณอย่างสมบูรณ์ อัลกอริทึมนี้ถูกนำเสนอครั้งแรกในปี 2022 โดยทีมนักวิจัยจากมหาวิทยาลัยสแตนฟอร์ดภายใต้การนำของ Tri Dao[1]

แนวคิดหลักของ FlashAttention คือการจัดระเบียบการคำนวณใหม่โดยคำนึงถึงลำดับชั้นของหน่วยความจำ GPU ซึ่งช่วยลดจำนวนการเข้าถึงหน่วยความจำที่ช้า และขจัดคอขวดหลักของกลไก attention แบบมาตรฐาน

ปัญหาของ Attention แบบมาตรฐาน

กลไก self-attention แบบมาตรฐานในสถาปัตยกรรม transformer คำนวณตามสูตร: Attention(Q,K,V)=softmax(QKTdk)V โดยที่ Q, K, V คือเมทริกซ์ของ query, key และ value

ปัญหาหลักของแนวทางนี้คือ ความซับซ้อนกำลังสอง ทั้งในด้านเวลาและหน่วยความจำ (O(N²)) เทียบกับความยาวลำดับ N[1] ในการนำไปใช้อย่างตรงไปตรงมา จำเป็นต้องคำนวณและจัดเก็บเมทริกซ์ attention เต็มรูปแบบ S ขนาด N×N ใน GPU ซึ่งนำไปสู่ปัญหาวิกฤตสองประการ:

  1. การใช้หน่วยความจำสูง: การจัดเก็บเมทริกซ์ N×N เป็นไปไม่ได้เมื่อทำงานกับ context ที่ยาว
  2. การดำเนินการ I/O: คอขวดหลักไม่ใช่จำนวนการดำเนินการทางคณิตศาสตร์ แต่เป็นการเข้าถึงหน่วยความจำที่ช้าของ GPU อย่างต่อเนื่อง

ลำดับชั้นหน่วยความจำของ GPU

เพื่อทำความเข้าใจปัญหา สิ่งสำคัญคือต้องแยกแยะหน่วยความจำสองประเภทใน GPU (ตัวอย่าง NVIDIA A100):

  • SRAM (หน่วยความจำแบบสถิต): หน่วยความจำในชิปที่รวดเร็วแต่มีขนาดเล็ก (~20 MB) พร้อม bandwidth ขนาดใหญ่ (สูงถึง 19 TB/s)
  • HBM (High Bandwidth Memory): หน่วยความจำขนาดใหญ่ที่ช้ากว่า (40–80 GB) พร้อม bandwidth ที่น้อยกว่ามาก (ประมาณ 1.5 TB/s)[2]

ความไม่สมมาตรนี้ทำให้อัลกอริทึม attention แบบมาตรฐาน ถูกจำกัดโดย memory bandwidth (memory-bound) เนื่องจากต้องอ่านและเขียนเมทริกซ์ขนาดใหญ่จาก HBM ที่ช้าอย่างต่อเนื่อง ซึ่งเป็นแหล่งหลักของความล่าช้า

นวัตกรรมหลักของ FlashAttention

FlashAttention เป็นอัลกอริทึมที่ ตระหนักถึง IO (IO-aware) ซึ่งแก้ปัญหาด้วยการลดการเข้าถึง HBM ให้น้อยที่สุด ทำได้โดยใช้เทคนิคหลักสามประการ

Tiling และการประมวลผลแบบบล็อก

แทนที่จะประมวลผลเมทริกซ์ทั้งหมดในครั้งเดียว FlashAttention แบ่งเมทริกซ์ input Q, K, V ออกเป็นบล็อกเล็กๆ (tile) ที่พอดีกับ SRAM ที่รวดเร็ว อัลกอริทึมจะโหลดบล็อกเหล่านี้ตามลำดับ ดำเนินการคำนวณ attention ทั้งหมดสำหรับบล็อกเหล่านั้น และอัปเดตผลลัพธ์สุดท้าย โดยไม่จัดเก็บเมทริกซ์ attention เต็มรูปแบบ ใน HBM ที่ช้า[1]

การคำนวณ Softmax แบบออนไลน์

ความก้าวหน้าทางเทคนิคที่สำคัญคือการคำนวณ Softmax แบบ "ออนไลน์" Softmax แบบมาตรฐานต้องการทราบองค์ประกอบทั้งหมดของเวกเตอร์ input เพื่อการ normalization FlashAttention ใช้อัลกอริทึมที่ดัดแปลงซึ่งช่วยให้คำนวณ Softmax เป็นส่วนๆ ได้ โดยรักษาค่ากลางสองค่า (ค่าสูงสุดปัจจุบันและผลรวมของ exponential) ที่อัปเดตเมื่อประมวลผลบล็อกใหม่ ทำให้ได้ผลลัพธ์ที่แม่นยำโดยไม่ต้องเข้าถึงเมทริกซ์ทั้งหมดพร้อมกัน[2]

การรวมการดำเนินการเป็น CUDA Kernel เดียว

การดำเนินการ attention ทั้งหมด (การคูณเมทริกซ์ QKᵀ, การ masking, Softmax, การคูณด้วย V) ถูกรวมเข้าเป็น CUDA kernel แบบรวม (fused kernel) เดียว ซึ่งลดจำนวนการดำเนินการอ่าน/เขียนใน HBM อย่างเด็ดขาด: แทนที่จะผ่านเมทริกซ์ทั้งหมดหลายรอบ อัลกอริทึมจะโหลดบล็อกเข้าสู่ SRAM ครั้งเดียว ดำเนินการคำนวณทั้งหมด และเขียนเฉพาะผลลัพธ์สุดท้าย

ประสิทธิภาพทางทฤษฎีและการปฏิบัติ

ความซับซ้อนและความเหมาะสมที่สุด

FlashAttention ลดการใช้หน่วยความจำจาก O(N²) เป็น O(N) ซึ่งให้การปรับขนาดเชิงเส้น มีการพิสูจน์แล้วว่าความซับซ้อน IO ของอัลกอริทึมนี้ เหมาะสมที่สุดในทางทฤษฎี สำหรับการคำนวณ attention ในลำดับชั้นหน่วยความจำสองระดับ กล่าวคือ ไม่สามารถคำนวณ attention แบบแม่นยำให้เร็วกว่านี้ได้โดยไม่เปลี่ยนแปลงฮาร์ดแวร์[3]

ผลลัพธ์เชิงประจักษ์

FlashAttention เวอร์ชันแรกแสดงให้เห็นการปรับปรุงที่มีนัยสำคัญ:

  • ความเร็ว:
    • BERT-large (ความยาว 512): เร็วขึ้น 15% ในการฝึก
    • GPT-2 (ความยาว 1K): เร็วขึ้น 3 เท่า
    • งาน Long-Range Arena (1K-4K): เร็วขึ้น 2.4 เท่า[1]
  • การประหยัดหน่วยความจำ: ประหยัดหน่วยความจำได้สูงถึง 20 เท่า เมื่อเทียบกับการนำไปใช้งานพื้นฐานแบบแม่นยำ
  • การปรับปรุงคุณภาพโมเดล: ด้วยความสามารถในการทำงานกับ context ที่ยาวขึ้น FlashAttention ไม่เพียงแต่ไม่สูญเสียคุณภาพของโมเดล แต่ยังปรับปรุงด้วย ตัวอย่างเช่น perplexity ของ GPT-2 ปรับปรุงขึ้น 0.7 จุด และความแม่นยำในงานจำแนกประเภทเอกสารยาวเพิ่มขึ้น 6.4 จุด[1]

วิวัฒนาการและการพัฒนาต่อเนื่อง

ความสำเร็จของ FlashAttention ได้จุดประกายอัลกอริทึมที่เน้นฮาร์ดแวร์ทั้งชุด

FlashAttention-2 (2023)

เวอร์ชันที่สองมุ่งเน้นการใช้ทรัพยากร GPU ให้เต็มประสิทธิภาพมากขึ้น ใน FlashAttention เดิมนั้น ประสิทธิภาพบน NVIDIA A100 อยู่ที่เพียง 25–40% ของค่าสูงสุด FlashAttention-2 นำเสนอการปรับปรุงใน parallelism การคำนวณ ซึ่งทำให้[4]:

  • บรรลุความเร็วที่ เพิ่มขึ้นสองเท่า เมื่อเทียบกับเวอร์ชันแรก
  • เพิ่มการใช้งาน GPU ได้ถึง 50–73% ของค่าสูงสุดทางทฤษฎี
  • ขยายการรองรับสำหรับ attention head ขนาด 256 รวมถึงสถาปัตยกรรม Multi-Query Attention (MQA)

FlashAttention-3 (2024)

เวอร์ชันที่สามได้รับการปรับให้เหมาะสมโดยเฉพาะสำหรับสถาปัตยกรรม GPU NVIDIA Hopper (H100)[5] โดยใช้ความสามารถฮาร์ดแวร์ใหม่ เช่น การทำงานแบบ asynchronous ของ Tensor Cores และการรองรับ FP8 ซึ่งทำให้:

  • บรรลุความเร็วที่ เพิ่มขึ้น 1.5–2 เท่า เมื่อเทียบกับ FlashAttention-2
  • บรรลุประสิทธิภาพสูงถึง 740 TFLOPS บน FP16 และใกล้เคียง 1.2 PFLOPS บน FP8

โซลูชันเฉพาะทาง

แนวคิดของ FlashAttention ได้รับการพัฒนาต่อในโปรเจกต์อื่นๆ:

  • FlashInfer (2025): เอนจิน attention ที่ปรับแต่งได้ ซึ่งเพิ่มประสิทธิภาพโดยเฉพาะสำหรับงาน inference ของ LLM โดยมุ่งเน้นการทำงานกับ KV-cache อย่างมีประสิทธิภาพในโหมดการสร้างแบบ streaming[6]
  • FlashMLA (2024): การนำ attention มาใช้พร้อมการบีบอัด context cache (latent attention) ซึ่งช่วยประหยัดหน่วยความจำสำหรับลำดับที่ยาวมากโดยสูญเสียข้อมูลน้อยที่สุด[7]

ผลกระทบต่ออุตสาหกรรมและระบบนิเวศ

FlashAttention กลายเป็นความก้าวหน้าพื้นฐานและกลายเป็น มาตรฐานของอุตสาหกรรม อย่างรวดเร็วสำหรับการฝึกและการ inference ของ LLM อย่างมีประสิทธิภาพ มันได้รับการรวมเข้ากับไลบรารีหลัก เช่น PyTorch และ Hugging Face และถูกใช้ในโมเดลภาษาขนาดใหญ่ส่วนใหญ่ (LLaMA, MPT, Falcon, Claude และอื่นๆ)

FlashAttention และเวอร์ชันต่อมามีบทบาทชี้ขาดในการขยาย หน้าต่าง context ของโมเดลภาษา: จาก 2–4 พัน token (GPT-3) ไปสู่ 128 พัน token (GPT-4) และแม้กระทั่งหลายล้าน token ในโมเดลเชิงทดลอง[8] อัลกอริทึมนี้ขจัดอุปสรรคหลักประการหนึ่งในการปรับขนาด transformer เปิดโอกาสใหม่สำหรับแอปพลิเคชัน AI ตั้งแต่การวิเคราะห์เอกสารยาวไปจนถึงความเข้าใจแบบ multimodal

ลิงก์

  • ที่เก็บโค้ดทางการของ FlashAttention บน GitHub

บรรณานุกรม

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