FlashAttention
FlashAttention to algorytm obliczania dokładnego attention, który ogranicza transfer danych między dużą pamięcią GPU, HBM, a znacznie mniejszą i szybszą pamięcią na układzie, SRAM. Zachowuje matematyczną operację Scaled Dot-Product Attention, lecz wykonuje ją blokami i nie zapisuje całej macierzy wyników porównań ani wag attention w HBM. Dao et al., §2–3.2, algorytm 1.
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.