ML Atlas

07 · Architektury · 4 min czytania · aktualizacja

Czym jest rekurencyjna sieć neuronowa (RNN) i dlaczego ma problem z długą pamięcią?

W skrócie

RNN czyta sekwencję krok po kroku i niesie stan ukryty, który streszcza dotychczasową historię. Prosta idea, ale pamięć szybko zanika z każdym krokiem.

Co to jest

Rekurencyjna sieć neuronowa (RNN, recurrent neural network) to sieć przetwarzająca sekwencję element po elemencie — słowo po słowie, pomiar po pomiarze — i przekazująca z kroku na krok stan ukryty h, czyli wektor streszczający wszystko, co do tej pory przeczytała. W każdym kroku nowy stan powstaje z poprzedniego stanu i bieżącego wejścia, zawsze tymi samymi wagami.

Intuicja: czytasz zdanie i nie zapamiętujesz każdego słowa osobno, tylko aktualizujesz w głowie „o czym to jest”. Po słowie „bank” stan może oznaczać „finanse albo ławka”, po słowie „kredyt” przesuwa się zdecydowanie w stronę finansów. RNN robi to samo, tylko jej „rozumienie” to kilkadziesiąt lub kilkaset liczb.

RNN były przez lata podstawowym narzędziem do tekstu, mowy, tłumaczenia maszynowego i szeregów czasowych. Klasyczną, prostą wersję opisał Elman (1990). W praktyce zastąpiły ją najpierw LSTM i GRU, a w przetwarzaniu języka — transformery.

Mechanizm — dlaczego tak działa

Równanie prostej RNN: h_t = tanh(W_h · h_(t−1) + W_x · x_t + b), a wyjście (np. przewidywane następne słowo) liczy się z h_t. Kluczowe jest to, że W_h, W_x i b są te same w każdym kroku. To współdzielenie wag w czasie — odpowiednik współdzielenia wag w przestrzeni w sieciach konwolucyjnych. Dzięki niemu sieć działa dla sekwencji dowolnej długości, a liczba parametrów nie zależy od długości tekstu.

Trening odbywa się przez propagację wsteczną w czasie (BPTT): sieć „rozwija się” w łańcuch tylu kopii, ile kroków ma sekwencja, i gradient płynie wstecz przez wszystkie te kopie. Tu leży słabość RNN. Gradient błędu w kroku t względem stanu sprzed k kroków jest iloczynem k czynników, każdy w przybliżeniu równy W_h razy pochodna tanh. Jeśli te czynniki są mniejsze od 1, gradient maleje wykładniczo (zanikający gradient); jeśli większe — rośnie wykładniczo (eksplodujący gradient).

Skutek zanikającego gradientu: sieć praktycznie nie może się nauczyć zależności oddalonych o kilkadziesiąt kroków, bo sygnał błędu nie dociera tak daleko. W zdaniu „Kot, którego widziałem wczoraj u sąsiadów na końcu ulicy, był rudy” trzeba połączyć „był rudy” z „kot” mimo kilkunastu słów przerwy. Bengio, Simard i Frasconi (1994) pokazali, że to trudność zasadnicza, a nie kwestia strojenia.

Eksplodujący gradient łatwiej opanować: obcina się normę gradientu (gradient clipping). Zanikający wymaga zmiany architektury — dlatego powstały LSTM i GRU z bramkami, które pozwalają informacji przepływać przez wiele kroków bez wielokrotnego mnożenia.

Druga słabość jest obliczeniowa: krok t wymaga wyniku kroku t − 1, więc RNN nie da się zrównoleglić po długości sekwencji. Na GPU, które lubią robić wszystko naraz, to duża wada — i jeden z powodów sukcesu transformerów.

Na przykładzie

Weźmy najprostszą RNN z jednym neuronem: h_t = tanh(0,5 · h_(t−1) + 1 · x_t), start h_0 = 0. Podajemy jeden impuls i potem same zera: x = 1, 0, 0, 0, 0. Stany: h_1 = tanh(1) ≈ 0,762, h_2 = tanh(0,381) ≈ 0,363, h_3 ≈ 0,180, h_4 ≈ 0,090, h_5 ≈ 0,045. Ślad impulsu mniej więcej połowi się z każdym krokiem — to cała „pamięć” tej sieci. Pochodna h_3 względem h_1 wynosi 0,5 · (1 − h_2²) · 0,5 · (1 − h_3²) ≈ 0,21; po 50 krokach czynnik rzędu 0,5^50 ≈ 9·10^(−16) nie przeniesie już żadnego sygnału uczenia.

Liczba parametrów nie zależy od długości sekwencji. Warstwa nn.RNN w PyTorch z wejściem 100-wymiarowym i stanem 128-wymiarowym ma 128·100 + 128·128 + 2·128 = 29 440 parametrów (PyTorch trzyma dwa wektory biasu). Te same 29 440 liczb obsłuży zdanie z 5 słów i dokument z 5000.

W praktyce

  • PyTorch: nn.RNN(input_size, hidden_size, batch_first=True); w praktyce prawie zawsze lepsze są nn.LSTM lub nn.GRU o tym samym interfejsie.
  • Zawsze stosuj obcinanie gradientu: torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0).
  • Sekwencje różnej długości pakuje się funkcją nn.utils.rnn.pack_padded_sequence, żeby sieć nie „czytała” sztucznych zer na końcu.
  • Wariant dwukierunkowy (bidirectional=True) czyta tekst w obie strony — dobry do klasyfikacji, niedozwolony przy generowaniu, gdy przyszłość jest nieznana.
  • Dla szeregów czasowych RNN warto porównać z prostymi modelami (np. regresją na opóźnieniach lub gradient boostingiem), które bywają równie dobre.

Najczęstsze pytania

Czym RNN różni się od zwykłej sieci gęstej?
Sieć gęsta dostaje całe wejście naraz i nie ma pamięci. RNN przetwarza sekwencję krok po kroku i przenosi stan ukryty, więc kolejność elementów ma znaczenie, a długość wejścia może być dowolna.
Dlaczego RNN zapomina początek długiego tekstu?
Bo gradient błędu cofa się przez każdy krok, mnożąc się po drodze przez czynniki zwykle mniejsze od 1. Po kilkudziesięciu krokach staje się praktycznie zerowy, więc sieć nie uczy się wykorzystywać dalekiej informacji.
Czy RNN są jeszcze używane?
Tak, ale niszowo: w urządzeniach o małej mocy, w przetwarzaniu strumieniowym i w niektórych modelach szeregów czasowych. Pojawiają się też nowe architektury rekurencyjne (np. modele przestrzeni stanów), które łączą liniowy koszt RNN z trenowaniem równoległym.

Źródła

  • Elman „Finding Structure in Time”, Cognitive Science 14(2), 1990.
  • Bengio, Simard, Frasconi „Learning Long-Term Dependencies with Gradient Descent is Difficult”, IEEE Transactions on Neural Networks 5(2), 1994.
  • Goodfellow, Bengio, Courville „Deep Learning”, MIT Press, 2016, rozdz. 10 („Sequence Modeling: Recurrent and Recursive Nets”).
  • Zhang i in. „Dive into Deep Learning”, d2l.ai, rozdz. 9 („Recurrent Neural Networks”).
  • Dokumentacja PyTorch: torch.nn.RNN, https://pytorch.org/docs/stable/generated/torch.nn.RNN.html

Zobacz też