ML Atlas

07 · Architektury · 5 min czytania · Interaktywne · aktualizacja

Czym jest samouwaga (self-attention) i jak tokeny „patrzą” na siebie nawzajem?

W skrócie

W samouwadze każdy token tej samej sekwencji tworzy zapytanie, klucz i wartość, a potem zbiera informację od pozostałych tokenów w proporcji do ich zgodności.

Co to jest

Samouwaga (self-attention) to mechanizm uwagi, w którym zapytania, klucze i wartości pochodzą z tej samej sekwencji. Każdy token zadaje pytanie wszystkim tokenom tego samego tekstu (łącznie z sobą), dostaje od nich wagi zgodności i zastępuje swoją reprezentację ważoną mieszanką ich wartości. Po jednej warstwie samouwagi wektor każdego słowa zawiera już informację o kontekście.

Intuicja: w zdaniu „Zamek był zamknięty na klucz” słowo „zamek” jest dwuznaczne — budowla czy mechanizm w drzwiach? Samouwaga pozwala mu „spojrzeć” na „zamknięty” i „klucz” i przesunąć swoją reprezentację w stronę zamka w drzwiach. W zdaniu „Zamek stał na wzgórzu” to samo słowo dostanie inną reprezentację, bo zbierze informację od innych sąsiadów. Tak powstają kontekstowe reprezentacje słów.

Samouwaga to serce transformera i wszystkich dużych modeli językowych. W modelach typu GPT jest przyczynowa (token widzi tylko poprzedników), w modelach typu BERT — dwukierunkowa (widzi cały tekst).

Mechanizm — dlaczego tak działa

Wejściem jest macierz X: n tokenów, każdy jako wektor o wymiarze d. Warstwa ma trzy macierze wag: W_Q, W_K, W_V. Liczy Q = X·W_Q, K = X·W_K, V = X·W_V, a następnie wynik softmax(Q·Kᵀ / √d_k) · V. Macierz Q·Kᵀ ma rozmiar n×n: element (i, j) mówi, jak bardzo token i jest zainteresowany tokenem j. Softmax działa wierszami, więc każdy token rozdziela swoją „uwagę” — sumę 1 — między wszystkie tokeny.

Dlaczego trzy różne projekcje, a nie po prostu iloczyn X·Xᵀ? Bo relacja „kto kogo potrzebuje” nie musi być symetryczna ani oparta na podobieństwie. Czasownik szuka podmiotu, zaimek — rzeczownika, do którego się odnosi. Osobne W_Q i W_K pozwalają nauczyć się, że token A pyta o cechę, którą token B oferuje, choć same w sobie A i B nie są podobne. W_V oddziela to, „po czym mnie znaleźć”, od tego, „co przekazuję dalej”.

Samouwaga różni się od RNN i splotu w dwóch ważnych punktach. Po pierwsze, każdy token ma bezpośredni dostęp do każdego innego — ścieżka między pierwszym a tysięcznym słowem ma długość 1, a nie 1000 kroków rekurencji. Po drugie, wszystkie pozycje liczy się równolegle jednym mnożeniem macierzy, co idealnie pasuje do GPU. To dwa główne powody, dla których transformery wyparły RNN.

Ceną jest koszt kwadratowy: macierz n×n trzeba policzyć w każdej warstwie i w każdej głowie. Podwojenie długości kontekstu to czterokrotnie więcej par. Dlatego tak dużo badań dotyczy wydajnych wariantów uwagi i dlatego okno kontekstu modeli jest ograniczone.

Ważna, mniej oczywista własność: samouwaga nie zna kolejności. Jeśli przestawisz tokeny wejścia, wyniki przestawią się dokładnie tak samo, ale ich wartości się nie zmienią (to tzw. ekwiwariancja względem permutacji). Dla samouwagi „pies gryzie człowieka” i „człowiek gryzie psa” to ten sam zbiór słów. Informację o kolejności trzeba dodać osobno — przez kodowanie pozycyjne.

Na przykładzie

Trzy tokeny A, B, C z wektorami X_A = [1, 0], X_B = [0, 1], X_C = [1, 1]. Projekcje: W_Q o wierszach [1, 0] i [0, 2], W_K o wierszach [0, 1] i [2, 0], W_V = macierz jednostkowa. Wychodzi Q_A = [1, 0], Q_B = [0, 2], Q_C = [1, 2] oraz K_A = [0, 1], K_B = [2, 0], K_C = [2, 1]. Wyniki Q·Kᵀ / √2 w wierszu C to [1,41; 1,41; 2,83], więc po softmaksie token C rozdziela uwagę: 0,16 na A, 0,16 na B i 0,67 na siebie. Token A ma wagi [0,11; 0,45; 0,45] — prawie nie patrzy na siebie, tylko na B i C. Wyniki (macierz wag razy V): A → [0,55; 0,89], B → [0,89; 0,55], C → [0,84; 0,84]. Przy kolejności C, A, B wyjście to dokładnie te same trzy wektory w nowej kolejności — sprawdziliśmy to numerycznie.

Koszt w liczbach: przy 1024 tokenach macierz uwagi ma 1 048 576 elementów na głowę i warstwę, przy 8192 tokenach — 67 108 864, czyli około 134 MB w formacie 16-bitowym, gdyby ją trzymać w pamięci w całości. Liczba wag warstwy nie zależy od długości tekstu: dla d = 768 (jak w GPT-2 small) cztery macierze projekcji (Q, K, V i wyjściowa) z biasami to 4·768² + 4·768 = 2 362 368 parametrów.

Ta ilustracja działa w przeglądarce z włączonym JavaScriptem: samouwaga w zdaniu „Kot zjadł rybę, bo był głodny” na ręcznie ustawionych wagach: „był” kieruje najwięcej uwagi na „Kot”.

W praktyce

  • PyTorch: nn.MultiheadAttention(d, num_heads, batch_first=True) i wywołanie attn(x, x, x) — trzy razy to samo wejście to właśnie samouwaga.
  • Dla modeli generujących: F.scaled_dot_product_attention(q, k, v, is_causal=True) albo jawna maska trójkątna.
  • Pamiętaj o kodowaniu pozycyjnym — bez niego model jest ślepy na kolejność słów.
  • Przy długich sekwencjach używa się wydajnych implementacji (FlashAttention), które nie materializują macierzy n×n w pamięci GPU, choć liczba operacji nadal rośnie kwadratowo.
  • Typowy błąd: pomylenie maski dopełnienia z maską przyczynową lub odwrócenie konwencji (True = „zablokuj” w PyTorch dla masek logicznych w nn.MultiheadAttention).

Najczęstsze pytania

Czym samouwaga różni się od zwykłej uwagi?
W klasycznej uwadze zapytania pochodzą z jednej sekwencji, a klucze i wartości z drugiej (np. tłumaczenie patrzy na zdanie źródłowe). W samouwadze wszystkie trzy pochodzą z tego samego tekstu, więc tokeny budują kontekst z siebie nawzajem.
Dlaczego koszt samouwagi rośnie kwadratowo?
Bo każdy z n tokenów liczy zgodność z każdym z n tokenów, co daje n² par. Przy dziesięciokrotnie dłuższym tekście par jest sto razy więcej.
Czy token patrzy też na samego siebie?
Tak, przekątna macierzy uwagi to uwaga tokenu na siebie. Często ma znaczną wagę, bo własna treść tokenu jest dla niego ważna, ale model może nauczyć się patrzeć głównie na innych.

Źródła

  • Vaswani i in. „Attention Is All You Need”, NeurIPS 2017, arXiv:1706.03762.
  • Devlin, Chang, Lee, Toutanova „BERT: Pre-training of Deep Bidirectional Transformers for Language Understanding”, NAACL 2019.
  • Zhang i in. „Dive into Deep Learning”, d2l.ai, rozdz. 11.6 („Self-Attention and Positional Encoding”).
  • Dao i in. „FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness”, NeurIPS 2022.
  • Dokumentacja PyTorch: torch.nn.MultiheadAttention, https://pytorch.org/docs/stable/generated/torch.nn.MultiheadAttention.html

Zobacz też