FlashAttention (PL)
FlashAttention — to rewolucyjny algorytm obliczania mechanizmu uwagi (attention), opracowany w celu znaczącego przyspieszenia trenowania i inferencji dużych modeli językowych (LLM) przy zachowaniu pełnej dokładności obliczeń. Algorytm został po raz pierwszy przedstawiony w 2022 roku przez zespół badaczy ze Uniwersytetu Stanforda pod kierownictwem Tri Dao[1].
Kluczowa idea FlashAttention polega na reorganizacji obliczeń z uwzględnieniem hierarchii pamięci GPU, co pozwala zminimalizować liczbę odwołań do wolnej pamięci i wyeliminować główne wąskie gardło standardowego mechanizmu uwagi.
Problematyka standardowego mechanizmu uwagi
Standardowy mechanizm samouwagi w transformerach jest obliczany według wzoru: gdzie Q, K, V to macierze zapytań, kluczy i wartości.
Głównym problemem tego podejścia jest złożoność kwadratowa pod względem czasu i pamięci (O(N²)) względem długości sekwencji N[1]. W naiwnej implementacji konieczne jest obliczanie i przechowywanie w pamięci GPU pełnej macierzy uwagi S o rozmiarze N×N, co prowadzi do dwóch krytycznych problemów:
- Duże zużycie pamięci: Przechowywanie macierzy N×N staje się niemożliwe przy pracy z długimi kontekstami.
- Operacje wejścia-wyjścia (IO): Głównym wąskim gardłem nie jest liczba operacji arytmetycznych, lecz ciągłe odwołania do wolnej pamięci GPU.
Hierarchia pamięci GPU
Dla zrozumienia problemu ważne jest rozróżnienie dwóch typów pamięci w GPU (na przykładzie NVIDIA A100):
- SRAM (pamięć statyczna): Szybka pamięć wewnątrzkrzemienna małej pojemności (~20 MB) o ogromnej przepustowości (do 19 TB/s).
- HBM (pamięć wysokiej przepustowości): Wolna pamięć dużej pojemności (40–80 GB) o znacznie mniejszej przepustowości (około 1,5 TB/s)[2].
Ta asymetria sprawia, że standardowy algorytm uwagi jest ograniczony przepustowością pamięci (memory-bound), ponieważ stale odczytuje i zapisuje duże macierze z wolnej pamięci HBM, co stanowi główne źródło opóźnień.
Kluczowe innowacje FlashAttention
FlashAttention jest algorytmem świadomym operacji IO (IO-aware), który rozwiązuje problem poprzez minimalizację odwołań do HBM. Osiąga się to za pomocą trzech głównych technik.
Tiling i przetwarzanie blokowe
Zamiast przetwarzać całą macierz naraz, FlashAttention dzieli wejściowe macierze Q, K, V na małe bloki (kafelki), które mieszczą się w szybkiej pamięci SRAM. Algorytm sekwencyjnie ładuje te bloki, wykonuje dla nich wszystkie obliczenia uwagi i aktualizuje końcowy wynik, nie zapisując pełnej macierzy uwagi w wolnej pamięci HBM[1].
Obliczanie Softmax w trybie online
Kluczowym przełomem technicznym stało się „online" obliczanie Softmax. Standardowy Softmax wymaga znajomości wszystkich elementów wektora wejściowego do normalizacji. FlashAttention wykorzystuje zmodyfikowany algorytm, który pozwala obliczać Softmax partiami. Przechowuje dwie wartości pośrednie (bieżące maksimum oraz sumę wykładników), które są aktualizowane w miarę przetwarzania kolejnych bloków, co pozwala uzyskać dokładny wynik bez dostępu do całej macierzy naraz[2].
Scalanie operacji w jedno jądro CUDA
Wszystkie operacje uwagi (mnożenie macierzowe QKᵀ, maskowanie, Softmax, mnożenie przez V) są łączone w jedno scalone jądro CUDA (fused kernel). Radykalnie zmniejsza to liczbę operacji odczytu/zapisu w HBM: zamiast wielokrotnych przebiegów po całej macierzy, algorytm ładuje blok do SRAM jeden raz, wykonuje wszystkie obliczenia i zapisuje wyłącznie końcowy wynik.
Efektywność teoretyczna i praktyczna
Złożoność i optymalność
FlashAttention redukuje zużycie pamięci z O(N²) do O(N), zapewniając liniowe skalowanie. Udowodniono, że złożoność IO algorytmu jest teoretycznie optymalna dla obliczania uwagi w dwupoziomowej hierarchii pamięci, co oznacza, że szybsze wykonanie dokładnej uwagi jest niemożliwe bez zmian sprzętowych[3].
Wyniki empiryczne
Pierwsza wersja FlashAttention wykazała znaczące ulepszenia:
- Przyspieszenie:
- BERT-large (długość 512): 15% przyspieszenia trenowania.
- GPT-2 (długość 1K): 3-krotne przyspieszenie.
- Zadania Long-Range Arena (1K–4K): 2,4-krotne przyspieszenie[1].
- Oszczędność pamięci: Do 20-krotnej oszczędności pamięci w porównaniu z dokładnymi implementacjami bazowymi.
- Poprawa jakości modeli: Dzięki możliwości pracy z dłuższymi kontekstami, FlashAttention nie tylko nie traci na jakości, ale wręcz ją poprawia. Na przykład perpleksja GPT-2 poprawiła się o 0,7 punktu, a dokładność w zadaniach klasyfikacji długich dokumentów wzrosła o 6,4 punktu[1].
Ewolucja i dalszy rozwój
Sukces FlashAttention zapoczątkował całą serię algorytmów zorientowanych sprzętowo.
FlashAttention-2 (2023)
Druga wersja była nakierowana na pełniejsze wykorzystanie zasobów GPU. W oryginalnym FlashAttention efektywność na NVIDIA A100 wynosiła zaledwie 25–40% maksimum. FlashAttention-2 wprowadziła usprawnienia w paralelizacji obliczeń, co pozwoliło[4]:
- Osiągnąć dwukrotne przyspieszenie w porównaniu z pierwszą wersją.
- Zwiększyć wykorzystanie GPU do 50–73% teoretycznego maksimum.
- Rozszerzyć wsparcie dla głowic uwagi o rozmiarze 256, a także dla architektur Multi-Query Attention (MQA).
FlashAttention-3 (2024)
Trzecia wersja została zoptymalizowana specjalnie dla architektury GPU NVIDIA Hopper (H100)[5]. Wykorzystuje nowe możliwości sprzętowe, takie jak asynchroniczność Tensor Cores i obsługa FP8, co pozwoliło:
- Osiągnąć kolejne 1,5–2-krotne przyspieszenie w porównaniu z FlashAttention-2.
- Osiągnąć wydajność do 740 TFLOPS na FP16 i blisko 1,2 PFLOPS na FP8.
Rozwiązania specjalizowane
Idee FlashAttention zostały rozwinięte w innych projektach:
- FlashInfer (2025): Konfigurowalny silnik uwagi, zoptymalizowany specjalnie do zadań inferencji LLM. Koncentruje się na efektywnej obsłudze pamięci podręcznej KV w trybie generowania strumieniowego[6].
- FlashMLA (2024): Implementacja uwagi ze kompresją pamięci podręcznej kontekstu (latent attention), pozwalająca oszczędzać pamięć na bardzo długich sekwencjach przy minimalnej utracie informacji[7].
Wpływ na przemysł i ekosystem
FlashAttention stał się fundamentalnym przełomem i szybko przekształcił się w standard branżowy dla efektywnego trenowania i inferencji LLM. Został zintegrowany z kluczowymi bibliotekami, takimi jak PyTorch i Hugging Face, i jest stosowany w większości dużych modeli językowych (LLaMA, MPT, Falcon, Claude i inne).
To właśnie FlashAttention i jego kolejne wersje odegrały decydującą rolę w zwiększaniu okien kontekstowych modeli językowych: z 2–4 tys. tokenów (GPT-3) do 128 tys. tokenów (GPT-4), a nawet do milionów tokenów w eksperymentalnych modelach[8]. Algorytm usunął jedną z głównych przeszkód na drodze do skalowania transformerów, otwierając nowe możliwości dla aplikacji AI — od analizy długich dokumentów po rozumienie multimodalne.
Odnośniki
- Oficjalne repozytorium FlashAttention na GitHub
Literatura
- 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.
Przypisy
- ↑ 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.0 2.1 Дао, Три, и др. «FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness». OpenReview. [2]
- ↑ «We're Training AI Twice as Fast This Year as Last». IEEE Spectrum. [3]
- ↑ Дао, Три. «FlashAttention-2». tridao.me. [4]
- ↑ «FlashAttention-3: Fast and Accurate Attention with Asynchrony and Low-precision». PyTorch Blog. [5]
- ↑ «[2501.01005] FlashInfer: Efficient and Customizable Attention Engine for LLM Inference Serving». arXiv. [6]
- ↑ «GitHub - deepseek-ai/FlashMLA: FlashMLA: Efficient MLA decoding kernels». GitHub. [7]
- ↑ «The Evolution of Flash Attention: Revolutionizing Transformer Efficiency». Medium. [8]