Cross-Attention
Jak działa
1. Sekwencja „pytająca" (np. dekoder) jest rzutowana liniowo na macierz zapytań Q. 2. Sekwencja „źródłowa" (np. wyjście enkodera) jest rzutowana na macierze kluczy K i wartości V. 3. Wyliczane są iloczyny skalarne Q·Kᵀ, skalowane przez 1/√d_k i przepuszczane przez softmax, dając wagi uwagi wskazujące, jak bardzo każda pozycja pytająca odwołuje się do każdej pozycji źródłowej. 4. Wagi te mnożą wartości V, dając kontekstowe reprezentacje. 5. W praktyce operacja jest wielogłowicowa (multi-head): Q, K, V dzieli się na h głów liczonych równolegle, a wyniki łączy i rzutuje przez macierz wyjściową. Kluczowe: Q ma inne źródło niż K i V, sekwencje mogą mieć różne długości, a maska przyczynowa nie jest stosowana (w przeciwieństwie do maskowanej self-attention dekodera). Podczas dekodowania autoregresyjnego K i V z enkodera liczone są raz i cache'owane dla wszystkich kroków.
Rozwiązany problem
Jak pozwolić jednej sekwencji (lub modalności) korzystać z informacji zawartej w innej, potencjalnie o innej długości? Cross-attention rozwiązuje problem warunkowego łączenia dwóch reprezentacji — np. dopasowania tłumaczenia do zdania źródłowego, obrazu do opisu tekstowego czy akcji robota do obserwacji wizualnej — bez konieczności kompresji sekwencji źródłowej do pojedynczego wektora, co było wąskim gardłem wcześniejszych architektur enkoder-dekoder.
Komponenty
Liniowa projekcja W_Q stosowana do sekwencji pytającej. To ona odróżnia cross-attention od self-attention — Q pochodzi z innego strumienia niż K i V.
Projekcje W_K i W_V stosowane do sekwencji źródłowej. W dekodowaniu autoregresyjnym liczone raz i cache'owane, bo źródło się nie zmienia.
softmax(Q·Kᵀ / √d_k)·V. Bez maski przyczynowej; opcjonalna maska dopełnienia (padding) po stronie źródła.
Oficjalna
Konkatenacja h głowic i liniowa projekcja W_O.
Implementacja
Q musi pochodzić z sekwencji pytającej (dekoder), a K i V z sekwencji źródłowej (enkoder). Zamiana źródeł niszczy warunkowanie.
Cross-attention nie używa maski przyczynowej — potrzebna jest tylko maska dopełnienia po stronie źródła. Dodanie maski causal błędnie ogranicza dostęp do wejścia.
W dekodowaniu autoregresyjnym K i V ze źródła nie zmieniają się; przeliczanie ich w każdym kroku marnuje obliczenia.
Ewolucja
Bahdanau i in. wprowadzają uwagę między dekoderem a stanami enkodera w tłumaczeniu maszynowym — dekoder odwołuje się do sekwencji źródłowej.
Vaswani i in. definiują warstwę encoder-decoder attention ze skalowaną uwagą iloczynu skalarnego i wieloma głowicami; Q z dekodera, K i V z enkodera.
Perceiver używa cross-attention, by rzutować bardzo duże wejścia na zwartą tablicę latentną, oddzielając koszt od długości wejścia.
Rombach i in. wprowadzają cross-attention w UNet, by warunkować generowanie obrazu na osadzeniach tekstu — podstawa Stable Diffusion.
Flamingo wstrzykuje reprezentacje wizualne do zamrożonego modelu językowego przez bramkowane warstwy cross-attention.
Hiperparametry (konfigurowalne osie)
Liczba równoległych głowic uwagi.
Wymiar na głowicę, używany w skalowaniu 1/√d_k.
Wymiar osadzeń wejściowych i wyjściowych warstwy.
Złożoność obliczeniowa
Złożoność czasowa: O(n_q · n_kv · d). Złożoność przestrzenna: O(n_q · n_kv).
Wąskie gardło obliczeniowe
Główny koszt to iloczyny skalarne między zapytaniami a kluczami oraz normalizacja softmax; przy długich sekwencjach źródłowych rośnie liniowo z n_kv.
Paradygmat wykonania
Uwaga gęsta: każda pozycja pytająca odwołuje się do wszystkich pozycji źródłowych (poza maskowanym dopełnieniem).
Równoległość
W treningu wszystkie pozycje pytające liczone równolegle. Przy dekodowaniu autoregresyjnym pozycje pytające powstają sekwencyjnie, ale K i V ze źródła liczone są raz i cache'owane.
Wymagania sprzętowe
Operacja zdominowana przez gęste iloczyny macierzowe (Q·Kᵀ, ·V), idealna dla rdzeni tensorowych.
Systoliczne jednostki mnożenia macierzy dobrze obsługują wsadowe iloczyny uwagi.
Działa na CPU i innych akceleratorach, lecz z gorszą przepustowością przy długich sekwencjach źródłowych.