FlashAttention-3 (VI)
FlashAttention-3 — đây là thuật toán để tối ưu hóa cơ chế attention trong các mạng nơ-ron transformer, được phát triển nhằm tận dụng tối đa các khả năng phần cứng của GPU kiến trúc NVIDIA Hopper (H100)[1]. Thuật toán được giới thiệu vào năm 2024 bởi một nhóm các nhà nghiên cứu từ các công ty Colfax Research, Meta, NVIDIA, Georgia Tech, Đại học Princeton và Together AI. Công trình đã được chấp nhận tại hội nghị NeurIPS 2024 và được ghi nhận là spotlight[2].
FlashAttention-3 là lần lặp thứ ba trong họ thuật toán, kế tiếp FlashAttention (2022) và FlashAttention-2 (2023). Mục tiêu chính của nó là tăng tốc đáng kể quá trình huấn luyện và suy luận (inference) của các mô hình ngôn ngữ lớn (LLM), đồng thời vẫn duy trì độ chính xác của các phép tính.
Giới thiệu và bối cảnh
Vấn đề của cơ chế attention
Thành phần cốt lõi của transformer là cơ chế tự chú ý (self-attention), tuy nhiên độ 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²)) khi độ dài chuỗi đầu vào (n) tăng lên[1]. Điều này tạo ra một "nút cổ chai" nghiêm trọng, vì các GPU hiện đại được tối ưu hóa cho các phép nhân ma trận nhanh, nhưng việc tính toán các hàm mũ (ví dụ như trong Softmax) lại chậm hơn nhiều bậc. Ngoài ra, trong cách triển khai thông thường, bộ nhớ GPU phải lưu trữ một tensor attention trung gian lớn, điều này hạn chế khả năng mở rộng của các mô hình.
FlashAttention và FlashAttention-2
Để giải quyết vấn đề này, vào năm 2022 FlashAttention đã được đề xuất, giúp giảm số lần truy cập vào bộ nhớ toàn cục (HBM) chậm nhờ hai kỹ thuật:
- Xử lý theo khối (tiling): Các phép tính được chia thành các khối (tile) và xử lý trong bộ nhớ on-chip nhanh (SRAM).
- Hợp nhất các phép tính: Tất cả các phép tính (nhân ma trận, Softmax) được thực hiện trong một kernel GPU duy nhất mà không ghi kết quả trung gian vào bộ nhớ toàn cục.
Điều này cho phép giảm độ phức tạp bộ nhớ từ bậc hai xuống tuyến tính và tăng tốc tính toán lên 2–4 lần.
Vào năm 2023, phiên bản cải tiến FlashAttention-2 được giới thiệu, tối ưu hóa việc song song hóa các phép tính. Trên GPU kiến trúc NVIDIA Ampere (A100), nó đạt được ~70% hiệu suất lý thuyết tối đa[3]. Tuy nhiên, trên kiến trúc mới hơn NVIDIA Hopper (H100), hiệu quả của nó thấp hơn đáng kể — khoảng 35%[1]. Nguyên nhân là thuật toán không tận dụng các khả năng phần cứng mới của Hopper, điều này trở thành động lực để tạo ra FlashAttention-3.
Các khả năng phần cứng mới của GPU Hopper (H100)
Kiến trúc NVIDIA Hopper cung cấp một số tính năng mới mà FlashAttention-3 sử dụng để đạt hiệu suất tối đa[4]:
- WGMMA (Warpgroup Matrix Multiply-Accumulate): Loại lệnh mới cho các tensor core, thực hiện phép nhân ma trận với mức tăng hiệu suất gần gấp đôi so với kiến trúc Ampere.
- TMA (Tensor Memory Accelerator): Mô-đun phần cứng giúp tăng tốc truyền dữ liệu giữa bộ nhớ toàn cục (HBM) và bộ nhớ dùng chung (shared memory). TMA tự động thực hiện các tính toán địa chỉ, giảm tải cho các đơn vị tính toán.
- Định dạng FP8: Hỗ trợ phần cứng cho định dạng dữ liệu dấu phẩy động 8-bit, giúp nhân đôi hiệu suất lý thuyết so với FP16, nhưng tiềm ẩn nguy cơ mất độ chính xác do phạm vi động hạn chế.
Các đổi mới kỹ thuật của FlashAttention-3
Thuật toán triển khai ba phương pháp tối ưu hóa chính, được phát triển đặc biệt cho kiến trúc Hopper[4]:
1. Thực thi bất đồng bộ và chuyên biệt hóa warp
FlashAttention-3 sử dụng nguyên tắc warp-specialization, trong đó các nhóm luồng (warps) khác nhau trên GPU chuyên biệt hóa cho các nhiệm vụ khác nhau:
- Producer warps: Tải dữ liệu từ bộ nhớ toàn cục bằng TMA.
- Consumer warps: Thực hiện phép nhân ma trận trên các tensor core.
Nhờ tính bất đồng bộ phần cứng của Hopper, các phép tính này được thực hiện song song theo thời gian. Trong khi một nhóm warp thực hiện tính toán, nhóm khác đồng thời tải dữ liệu cho khối tiếp theo. Cách tiếp cận pipeline này, được tổ chức theo nguyên tắc "ping-pong" (ping-pong scheduling), cho phép ẩn đi độ trễ từ các phép tính chậm (ví dụ: Softmax) và tận dụng tối đa tất cả các mô-đun chức năng của GPU.
2. Giảm thiểu các phép tính bộ nhớ
Thuật toán giữ nguyên triết lý tiling từ các phiên bản trước, nhưng tích cực sử dụng TMA để tải bất đồng bộ các khối dữ liệu tiếp theo song song với các phép tính hiện tại. Việc truyền dữ liệu từ HBM chậm sang SRAM nhanh thực tế được thực hiện "trong bóng tối" của các phép tính chính, nhờ đó GPU ít bị chờ đợi dữ liệu hơn.
3. Độ chính xác thấp (FP8) với giảm thiểu lỗi lượng tử hóa
Chuyển sang FP8 giúp tăng gấp đôi tốc độ, nhưng có thể dẫn đến mất độ chính xác đáng kể do lượng tử hóa. Để giải quyết vấn đề này, các nhà phát triển đã triển khai phương pháp incoherent processing[4]. Nội dung của phương pháp như sau:
- Trước khi tính toán attention, các vector đặc trưng (query Q và key K) được nhân với một ma trận trực giao ngẫu nhiên (ví dụ: ma trận Hadamard).
- Phép biến đổi này "trải đều" các giá trị có mô-đun bất thường lớn (outlier) ra tất cả các chiều, cân bằng phân phối của chúng.
- Sau đó thực hiện lượng tử hóa sang FP8, lúc này diễn ra với lỗi nhỏ hơn.
- Vì phép biến đổi là trực giao, nó không làm méo kết quả cuối cùng của attention (QKᵀ), vì hiệu ứng của ma trận được triệt tiêu khi nhân.
Kỹ thuật này cho phép giảm lỗi tính toán attention trong FP8 khoảng 2,6 lần so với việc áp dụng FP8 tiêu chuẩn mà không có phép biến đổi[4].
Hiệu suất và ý nghĩa
Việc áp dụng các kỹ thuật kể trên đã giúp FlashAttention-3 đạt được sự vượt trội đáng kể so với các phiên bản trước trên GPU H100:
- Tăng tốc 1,5–2 lần so với FlashAttention-2.
- Hiệu suất sử dụng GPU cao: Đạt ~75–85% hiệu suất lý thuyết tối đa của H100.
- Thông lượng:
- Lên đến 740–840 TFLOPS với độ chính xác nửa (FP16/BF16).
- Lên đến 1,2–1,3 PFLOPS (petaflops) khi sử dụng độ chính xác 8-bit (FP8)[2].
Hiệu quả cao của FlashAttention-3 ảnh hưởng trực tiếp đến việc phát triển và ứng dụng LLM:
- Giảm thời gian huấn luyện: Tăng tốc attention 75–100% giúp rút ngắn đáng kể thời gian huấn luyện mô hình, vốn có thể kéo dài hàng tuần hoặc hàng tháng.
- Tăng cửa sổ ngữ cảnh: Các mô hình có thể xử lý hiệu quả các chuỗi dài hơn (hàng trăm nghìn token), điều này quan trọng cho việc phân tích các tài liệu hoặc mã nguồn lớn[1].
- Sử dụng tài nguyên hợp lý: Cho phép đạt được hiệu suất tương đương với số lượng GPU ít hơn hoặc đạt tốc độ cao hơn trên cùng một thiết bị, giúp giảm chi phí triển khai mô hình.
Khả năng tiếp cận và tích hợp
Các tác giả đã công bố mã nguồn của FlashAttention-3 theo giấy phép mã nguồn mở trên GitHub[4]. Dự kiến nó sẽ được tích hợp vào các framework học sâu hàng đầu như PyTorch và thư viện Hugging Face Transformers, giúp công nghệ này trở nên dễ tiếp cận hơn với đông đảo các nhà phát triển và nhà nghiên cứu. Các phiên bản trước đã trở thành tiêu chuẩn thực tế trong ngành, và FlashAttention-3 nhiều khả năng sẽ tiếp tục xu hướng này.
Liên kết
- Kho lưu trữ chính thức của FlashAttention trên GitHub
- Blog của Together AI với thông báo về FlashAttention-3
Tài liệu tham khảo
- 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.
Ghi chú
- ↑ 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]