FlashAttention-2 (VI)
FlashAttention-2 — đây là một thuật toán cải tiến được thiết kế để tính toán cơ chế attention trong các mô hình ngôn ngữ lớn (LLM). Thuật toán được phát triển bởi Tri Dao và các nhà nghiên cứu từ Đại học Stanford, được giới thiệu vào tháng 7 năm 2023[1]. Mục tiêu chính của nó là tăng tốc đáng kể quá trình huấn luyện và inference (suy luận) của các mô hình transformer thông qua việc sử dụng tài nguyên phần cứng GPU hiệu quả hơn, đồng thời duy trì sự đồng nhất hoàn toàn trong tính toán so với cơ chế attention tiêu chuẩn, tức là không mất độ chính xác.
FlashAttention-2 là sự tiếp nối logic của thuật toán FlashAttention được giới thiệu bởi cùng nhóm nghiên cứu vào năm 2022. Phiên bản mới giải quyết vấn đề tải GPU không đầy đủ vốn tồn tại ở phiên bản tiền nhiệm, và đạt được tốc độ tăng gần gấp đôi so với phiên bản đầu tiên.
Tiền đề: vấn đề attention trong transformer
Cơ chế self-attention tiêu chuẩn là điểm nghẽn cổ chai khi làm việc với các chuỗi văn bản dài trong transformer. Độ phức tạp tính toán và mức tiêu thụ bộ nhớ của nó tăng theo bậc hai (O(N²)) tùy thuộc vào độ dài chuỗi (N), điều này đặt ra những hạn chế nghiêm trọng đối với độ dài ngữ cảnh tối đa và khả năng mở rộng của LLM[1].
Để giải quyết vấn đề này, thuật toán FlashAttention đã được giới thiệu vào năm 2022[2]. Các ý tưởng chủ đạo của nó:
- Tính đến phân cấp bộ nhớ GPU (IO-awareness): Thuật toán giảm thiểu các thao tác đọc/ghi tốn kém giữa bộ nhớ GPU chậm (HBM) và bộ nhớ tĩnh nhanh (SRAM) trên chip.
- Xử lý theo khối (tiling): Các phép tính được chia thành các khối nhỏ (tile) và được xử lý trong SRAM nhanh, giúp tránh việc vật chất hóa toàn bộ ma trận attention trong bộ nhớ.
Điều này cho phép đạt được mức tiêu thụ bộ nhớ tăng theo tuyến tính (O(N)) và tăng tốc 2–4 lần so với các triển khai tiêu chuẩn[2]. FlashAttention đã được sử dụng rộng rãi và góp phần vào sự xuất hiện của các mô hình có ngữ cảnh tăng đáng kể, ví dụ từ 2–4 nghìn token (GPT-3) lên 128 nghìn (GPT-4) và hơn thế nữa[3]. Chẳng hạn, trong mô hình Falcon-40B, việc sử dụng FlashAttention đã tăng tốc inference lên 3 lần, và hiệu suất tạo văn bản tổng thể lên 5 lần so với GPT-3[4].
Phát triển và mục tiêu của FlashAttention-2
Mặc dù thành công, phiên bản đầu tiên của FlashAttention chưa sử dụng hết tài nguyên tính toán của GPU. Trên card đồ họa NVIDIA A100, hiệu suất chỉ đạt 25–40% so với mức tối đa lý thuyết (FLOPs/s)[1]. Nguyên nhân chính là do tải Streaming Multiprocessors không tối ưu và các thao tác thừa với bộ nhớ dùng chung[5].
Mục tiêu của FlashAttention-2 là tiếp tục tăng tốc tính toán thông qua song song hóa công việc hiệu quả hơn và giảm thiểu các thao tác phụ trợ. Thuật toán đã được viết lại hoàn toàn bằng cách sử dụng các primitive cấp thấp của thư viện NVIDIA CUTLASS 3.x để đạt hiệu suất tối đa[6].
Kiến trúc kỹ thuật và nguyên lý hoạt động
FlashAttention-2 giới thiệu ba cải tiến chính để tăng tính song song và hiệu quả[1]:
1. Giảm thiểu các thao tác không phải ma trận
Thuật toán giảm số lượng các thao tác dấu phẩy động phụ trợ không phải là phép nhân ma trận (non-matmul FLOPs). Vì các tensor core của GPU được tối ưu hóa đặc biệt cho các phép toán ma trận (GEMM) và thực hiện chúng nhanh hơn tới 16 lần, thay đổi này cho phép sử dụng các khối GPU hiệu quả nhất trong phần lớn thời gian.
2. Tính song song được cải thiện
Trong FlashAttention gốc, công việc trên một "đầu" attention không được song song hóa, dẫn đến thời gian chờ với các chuỗi dài và kích thước batch nhỏ. FlashAttention-2 giới thiệu song song hóa liên khối: giờ đây các phép tính cho một đầu attention được phân phối giữa các Streaming Multiprocessors khác nhau của GPU, điều này làm tăng đáng kể mức độ tải của chúng.
3. Phân chia công việc được tối ưu hóa trong một khối
Ở cấp độ một khối tính toán, công việc đã được phân phối lại giữa các nhóm luồng (warp) để giảm trao đổi dữ liệu qua bộ nhớ dùng chung (shared memory). Điều này làm giảm số lượng thao tác đọc/ghi thừa cần thiết cho chuẩn hóa Softmax.
Hiệu suất và hiệu quả
Nhờ các cải tiến kiến trúc, FlashAttention-2 thể hiện sự tăng hiệu suất đáng kể:
- Tăng tốc gấp đôi: Thuật toán hoạt động nhanh hơn khoảng 2 lần so với phiên bản FlashAttention đầu tiên[1].
- Tận dụng GPU cao: Trên GPU NVIDIA A100 đạt 50–73% so với thông lượng tối đa lý thuyết (TFLOPs), gần với hiệu quả của các thao tác nhân ma trận được tối ưu hóa (GEMM)[1].
- Tốc độ tính toán kỷ lục:
Sự tăng hiệu suất như vậy cho phép, ví dụ, huấn luyện mô hình với cửa sổ ngữ cảnh 16k token trong cùng khoảng thời gian trước đây cần cho cửa sổ 8k token[5]. Điều quan trọng là thuật toán vẫn chính xác và mang tính xác định, do đó việc áp dụng nó không ảnh hưởng đến chất lượng dự đoán của mô hình[8].
Ứng dụng và tích hợp vào hệ sinh thái
FlashAttention-2 nhanh chóng trở thành công cụ tiêu chuẩn trong hệ sinh thái LLM. Nó được tích hợp vào nhiều framework và thư viện phổ biến:
- PyTorch: Hỗ trợ gốc.
- Hugging Face Transformers: Hỗ trợ được bật bằng tham số `attn_implementation="flash_attention_2"` khi tải mô hình[9]. Tương thích với hàng chục kiến trúc (GPT, Llama, Falcon, BERT và các kiến trúc khác)[10].
- TensorRT-LLM, xFormers và Triton: Thuật toán được triển khai cho các nền tảng này, đảm bảo ứng dụng rộng rãi[7].
Việc tích hợp cho phép dễ dàng kết hợp FlashAttention-2 với các phương pháp tối ưu hóa khác như lượng tử hóa (GPTQ, QLoRA) và fine-tuning hiệu quả (PEFT)[9].
So sánh với các phiên bản kế tiếp
FlashAttention-3
Nghiên cứu trong lĩnh vực tối ưu hóa attention vẫn tiếp tục. Vào tháng 7 năm 2024, Tri Dao đã giới thiệu FlashAttention-3, nhắm đến việc tận dụng khả năng của kiến trúc GPU NVIDIA Hopper (H100/H200). Các tính năng mới chính[3]:
- Hỗ trợ FP8: Sử dụng tính toán dấu phẩy động 8-bit để tăng tốc thêm.
- Các thao tác bất đồng bộ: Sử dụng các khả năng bất đồng bộ của GPU hiệu quả hơn.
FlashAttention-3 cung cấp tốc độ tăng 1,5–2 lần so với FlashAttention-2 trên GPU H100, đạt hiệu suất lên tới 740 TFLOP/s (75% so với mức tối đa lý thuyết)[11].
Tài liệu tham khảo
- 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.
Chú thích
- ↑ 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.0 2.1 «Optimizing LLMs for Speed and Memory». Hugging Face Documentation. [2]
- ↑ 3.0 3.1 Дао, Три. «FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision». Tri Dao's Blog. [3]
- ↑ «FlashAttention vs FlashAttention-2 - an Analysis». E2E Networks Blog. [4]
- ↑ 5.0 5.1 «FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning». OpenReview. [5]
- ↑ «FlashAttention-2». Hazy Research, Stanford University. [6]
- ↑ 7.0 7.1 7.2 Дао, Три. «FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning» (PDF). arXiv:2307.08691. [7]
- ↑ Рашка, Себастьян. «Llama 2 and FlashAttention 2». Ahead of AI Magazine. [8]
- ↑ 9.0 9.1 Белькада, Юнес. «Faster and more memory efficient models with Flash Attention 2!». LinkedIn. [9]
- ↑ «GPU inference». Hugging Face Documentation. [10]
- ↑ Дао, Три, и др. «FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision». arXiv:2407.08608 [cs.LG], 11 июля 2024 г. [11]