FlashAttention
Attention porównuje wiele par fragmentów tekstu. Gdy zapisujemy wszystkie pośrednie wyniki do dużej pamięci karty obliczeniowej, przesyłanie danych może stać się kosztowne. Chcesz wykonać ten sam rachunek, ograniczając przenoszenie i przechowywanie wielkiej tabeli.
FlashAttention organizuje obliczanie attention blokami w małej, szybkiej pamięci układu. Nie zapisuje całej macierzy porównań i wag do dużej pamięci GPU, czyli karty używanej do równoległych obliczeń.
Każda porcja wnosi część wyniku, a algorytm poprawnie aktualizuje normalizację i sumę. Nie można po prostu policzyć osobnych średnich porcji i ich uśrednić: porcje mogą mieć różne udziały.
Zachowana jest matematyczna operacja dokładnego attention, z typowymi różnicami arytmetyki komputerowej. Dla gęstego attention liczba porównań nadal rośnie kwadratowo. Zysk wynika z organizacji pamięci i wykonania, a nie usunięcia zależności między parami.
Źródło mechanizmu: Dao et al., §2–3.2, algorytm 1.
Mechanizm i szczegóły

Dla gęstego attention liczba działań nadal rośnie kwadratowo z długością sekwencji. Zysk pochodzi z organizacji obliczeń i pamięci. Podczas treningu część wielkości jest ponownie liczona w backward zamiast przechowywana. „Dokładne” oznacza brak przybliżania samej operacji; inna kolejność działań zmiennoprzecinkowych może dać drobne różnice numeryczne.
Nie wystarczy uśrednić wyników bloków
Softmax normalizuje wagi względem wszystkich dozwolonych kluczy. Bloki mogą mieć bardzo różne sumy nienormalizowanych wag. Każdemu trzeba więc nadać właściwy udział w wyniku.
Można to robić stopniowo, utrzymując maksimum score , sumę i licznik . Wynik wynosi . Gdy nowy blok zwiększa maksimum do , wcześniejsze i trzeba przeskalować przez przed dodaniem nowych składników. Mechanizm stabilnej normalizacji wyjaśniają Milakov i Gimelshein, §3, algorytm 3.
Własny przykład używa gotowych scores oraz values . Wagi przed normalizacją mają proporcje , więc wynik to . Wyniki dwóch bloków to 0 i 1. Ich zwykła średnia daje błędne 0,5.
Demonstracja izoluje agregację dla jednej query i skalarnych values. Pełny kernel liczy bloki scores z Q i K na GPU; tutaj gotowe liczby pozwalają sprawdzić rachunek ręcznie. Przełączenie wielkości bloku nie jest benchmarkiem.
Grouped-Query Attention ogranicza liczbę głów K/V, a FlashAttention sposób przesyłania i przetwarzania danych. Te techniki można łączyć. FlashAttention 2 jest późniejszym rozwinięciem implementacji; opisany tu mechanizm pochodzi z pierwszej pracy.
Wykorzystuję treści generowane przez AI jako część mojego codziennego procesu nauki.