07 · Architektury · 4 min czytania · Interaktywne · aktualizacja
Czym jest mechanizm uwagi (attention) w sieciach neuronowych i jak działa?
W skrócie
Uwaga pozwala modelowi w każdym kroku sięgnąć do dowolnego fragmentu wejścia i wziąć z niego tyle, ile pasuje do bieżącego pytania — średnią ważoną zgodnością.
Co to jest
Mechanizm uwagi (attention) to operacja, w której model dla bieżącego „pytania” (zapytania, query) ocenia, jak bardzo pasuje do niego każdy z dostępnych elementów (kluczy, keys), zamienia te oceny na wagi sumujące się do 1 i zwraca średnią ważoną skojarzonych z nimi treści (wartości, values). Model nie musi więc ściskać całego wejścia w jeden wektor — w każdym kroku może „spojrzeć” tam, gdzie akurat jest potrzebna informacja.
Intuicja: to miękkie wyszukiwanie w słowniku. Zwykły słownik zwraca wartość dla klucza, który dokładnie pasuje. Uwaga zwraca mieszankę wszystkich wartości, w której więcej jest tych, których klucze są podobne do zapytania. Tłumacząc zdanie „Kot pije mleko” na angielski, przy generowaniu słowa „milk” model kładzie największą wagę na „mleko”, choć technicznie widzi wszystkie słowa.
Mechanizm wprowadzili Bahdanau, Cho i Bengio (2015) do tłumaczenia maszynowego opartego na RNN. W 2017 roku Vaswani i współpracownicy pokazali, że można zbudować cały model tylko z uwagi — tak powstał transformer, podstawa współczesnych modeli językowych.
Mechanizm — dlaczego tak działa
Najpopularniejsza wersja to skalowana uwaga iloczynowa (scaled dot-product attention): Attention(Q, K, V) = softmax(Q·Kᵀ / √d_k) · V. Krok po kroku: (1) iloczyn skalarny zapytania z każdym kluczem mierzy ich zgodność — duży, gdy wektory wskazują podobny kierunek; (2) dzielenie przez √d_k utrzymuje wyniki w rozsądnej skali; (3) softmax zamienia wyniki na dodatnie wagi sumujące się do 1; (4) wynik to suma wartości ważona tymi wagami.
Dlaczego to rozwiązuje problem RNN? Klasyczny koder–dekoder ściskał całe zdanie źródłowe w jeden wektor stanu, który dla długich zdań stawał się wąskim gardłem — jakość tłumaczenia spadała wraz z długością. Uwaga daje dekoderowi bezpośredni dostęp do stanu każdego słowa wejściowego. Ścieżka między dowolnymi dwoma pozycjami ma długość 1, więc gradient nie musi przechodzić przez dziesiątki kroków.
Dlaczego dzielić przez √d_k? Jeśli składowe wektorów q i k są losowe o średniej 0 i wariancji 1, iloczyn skalarny ma wariancję d_k, czyli odchylenie √d_k. Przy d_k = 64 wyniki rozrzucają się więc z odchyleniem około 8. Softmax z tak dużych liczb daje niemal zero-jedynkowe wagi, a wtedy jego gradienty są bliskie zeru i trening stoi. Dzielenie przez √d_k przywraca odchylenie około 1.
Wagi uwagi są zależne od danych — liczone na nowo dla każdego wejścia. To odróżnia uwagę od zwykłej warstwy gęstej, w której wagi połączeń są stałe po treningu. Uwaga jest też różniczkowalna, więc model uczy się propagacją wsteczną, jak formułować zapytania i klucze.
Zastrzeżenia: koszt liczenia wszystkich par zapytanie–klucz rośnie kwadratowo z długością sekwencji. Wagi uwagi bywają kuszące jako „wyjaśnienie” decyzji modelu, ale badania pokazują, że nie zawsze wskazują, co naprawdę wpłynęło na wynik — należy je traktować ostrożnie.
Na przykładzie
Trzy tokeny, wektory dwuwymiarowe (d_k = 2). Zapytanie q = [2, 1]. Klucze: k_1 = [1, 0], k_2 = [1, 1], k_3 = [0, −1]. Wartości: v_1 = [10, 0], v_2 = [0, 10], v_3 = [5, 5]. Iloczyny skalarne q·k: 2, 3 i −1. Po podzieleniu przez √2 ≈ 1,414: 1,41, 2,12 i −0,71. Softmax daje wagi 0,32, 0,64 i 0,04 (suma 1). Wynik: 0,32·[10, 0] + 0,64·[0, 10] + 0,04·[5, 5] ≈ [3,37; 6,63]. Zapytanie „patrzy” głównie na drugi token, trochę na pierwszy, a trzeci — o przeciwnym kierunku klucza — prawie pomija.
Skalowanie ma znaczenie. Gdyby te same wyniki były cztery razy większe (8, 12, −4), jak bywa przy wyższym wymiarze bez dzielenia przez √d_k, softmax dałby wagi 0,02, 0,98 i 0,00 — prawie twardy wybór jednego tokenu i niemal zerowy gradient dla pozostałych. Sprawdziliśmy też regułę skali: dla 100 tys. par losowych wektorów 64-wymiarowych o składowych z rozkładu N(0, 1) odchylenie standardowe iloczynu skalarnego wyniosło 8,0, a po podzieleniu przez √64 = 8 — 1,0.
W praktyce
- PyTorch:
torch.nn.functional.scaled_dot_product_attention(q, k, v, is_causal=True)— wydajna implementacja (m.in. FlashAttention) z maską przyczynową dla modeli generujących. - Gotowa warstwa z projekcjami i wieloma głowami:
nn.MultiheadAttention(embed_dim, num_heads, batch_first=True). - Maska przyczynowa zabrania patrzeć na przyszłe tokeny — bez niej model językowy „ściąga” odpowiedź podczas treningu.
- Maska dopełnienia (
key_padding_mask) wyklucza sztuczne tokeny dopełniające z wag uwagi. - Wizualizacja wag uwagi (mapy cieplne) pomaga w debugowaniu, ale nie jest pełnym wyjaśnieniem decyzji modelu.
Najczęstsze pytania
- Czym są zapytania, klucze i wartości?
- To trzy role tych samych danych. Zapytanie opisuje, czego szukamy; klucz — czym element „się reklamuje”; wartość — jaką treść odda, jeśli zostanie wybrany. W modelach są to zwykle trzy różne liniowe projekcje tych samych wektorów.
- Czym uwaga różni się od samouwagi?
- W uwadze krzyżowej zapytania pochodzą z jednej sekwencji (np. tłumaczenia), a klucze i wartości z innej (zdania źródłowego). W samouwadze wszystkie trzy pochodzą z tej samej sekwencji, więc każdy token patrzy na pozostałe tokeny tego samego tekstu.
- Czy wagi uwagi wyjaśniają, dlaczego model podjął decyzję?
- Tylko częściowo. Pokazują, skąd model pobrał informację w danej warstwie, ale model ma wiele warstw i głów, a inne konfiguracje wag mogą prowadzić do tej samej odpowiedzi. Do wyjaśniania lepiej używać metod atrybucji.
Źródła
- Bahdanau, Cho, Bengio „Neural Machine Translation by Jointly Learning to Align and Translate”, ICLR 2015, arXiv:1409.0473.
- Vaswani i in. „Attention Is All You Need”, NeurIPS 2017, arXiv:1706.03762.
- Jain, Wallace „Attention is not Explanation”, NAACL 2019.
- Zhang i in. „Dive into Deep Learning”, d2l.ai, rozdz. 11 („Attention Mechanisms and Transformers”).
- Dokumentacja PyTorch:
torch.nn.functional.scaled_dot_product_attention, https://pytorch.org/docs/stable/generated/torch.nn.functional.scaled_dot_product_attention.html