07 · Architektury · 5 min czytania · aktualizacja
Czym są sieci LSTM i GRU i jak bramki pomagają im pamiętać długo?
W skrócie
LSTM i GRU to sieci rekurencyjne z bramkami, które decydują, co zapomnieć, co zapisać i co odczytać. Dzięki temu pamiętają informację przez setki kroków.
Co to jest
LSTM (long short-term memory) i GRU (gated recurrent unit) to odmiany rekurencyjnych sieci neuronowych, w których przepływem informacji sterują bramki — małe warstwy z sigmoidą, zwracające liczby od 0 do 1. Bramka mnoży sygnał: 0 oznacza „zablokuj”, 1 — „przepuść w całości”. LSTM zaproponowali Hochreiter i Schmidhuber w 1997 roku; GRU to uproszczona wersja Cho i współpracowników z 2014 roku.
Intuicja: zwykła RNN to notatnik, który przy każdym nowym słowie jest w całości przepisywany od nowa — szybko gubią się w nim stare notatki. LSTM to notatnik z trzema decyzjami na każdym kroku: co wymazać (bramka zapominania), co dopisać (bramka wejścia) i co z notatek pokazać na zewnątrz (bramka wyjścia). Jeśli nic nie trzeba zmieniać, notatka przechodzi przez wiele kroków nietknięta.
W latach 2010–2017 LSTM były standardem w rozpoznawaniu mowy, tłumaczeniu maszynowym, rozpoznawaniu pisma i modelowaniu języka. W przetwarzaniu tekstu wyparły je transformery, ale LSTM i GRU wciąż sprawdzają się w szeregach czasowych i na słabszym sprzęcie.
Mechanizm — dlaczego tak działa
LSTM ma dwa stany: ukryty h (to, co sieć „pokazuje”) i stan komórki c (pamięć długotrwała). W każdym kroku z h_(t−1) i x_t liczy cztery wektory: bramkę zapominania f, bramkę wejścia i, bramkę wyjścia o (wszystkie przez sigmoidę) oraz kandydata na nową treść g (przez tanh). Aktualizacja: c_t = f ⊙ c_(t−1) + i ⊙ g, a następnie h_t = o ⊙ tanh(c_t), gdzie ⊙ oznacza mnożenie element po elemencie.
Kluczem jest dodawanie w równaniu na c_t. W zwykłej RNN nowy stan to nieliniowa funkcja starego, więc gradient przez każdy krok mnoży się przez macierz wag i pochodną tanh — i zanika. W LSTM pochodna c_t po c_(t−1) to po prostu f. Gdy sieć ustawi f blisko 1, pamięć i gradient przechodzą przez wiele kroków niemal bez strat. Hochreiter i Schmidhuber nazwali to „karuzelą stałego błędu” (constant error carousel). To ta sama idea, która później pojawi się w połączeniach rezydualnych: ścieżka addytywna zamiast iloczynu wielu czynników.
Bramki są zależne od danych. Sieć sama uczy się, kiedy zapominać — np. zerować informację o podmiocie, gdy zaczyna się nowe zdanie — a kiedy trzymać informację długo. Oryginalne LSTM z 1997 roku nie miało bramki zapominania; dodali ją Gers, Schmidhuber i Cummins w 2000 roku i od tej pory jest standardem.
GRU upraszcza ten układ: łączy stan komórki i stan ukryty w jeden, a zamiast trzech bramek ma dwie. Bramka aktualizacji z decyduje, jaką część starego stanu zachować, a jaką zastąpić nową treścią: h_t = (1 − z) ⊙ h_(t−1) + z ⊙ h̃_t (w oryginalnej pracy Cho i in. oraz w PyTorch role z i 1 − z są zamienione — to ta sama sieć). Bramka resetu r decyduje, ile przeszłości użyć przy liczeniu kandydata h̃_t. GRU ma około 25% mniej parametrów niż LSTM i w wielu zadaniach działa porównywalnie; żadna z nich nie wygrywa zawsze.
Ograniczenia: bramki łagodzą zanikanie gradientu, ale nie dają nieograniczonej pamięci — cała historia nadal musi zmieścić się w jednym wektorze o stałym rozmiarze. Obliczenia pozostają sekwencyjne, więc trening na długich tekstach jest wolny. Te dwa problemy rozwiązał mechanizm uwagi.
Na przykładzie
Jeden krok LSTM z liczbami (jeden wymiar). Komórka pamięta c_(t−1) = 1. Sieć „uznała”, że informację warto zachować: f = σ(3) ≈ 0,953, a nowego wiele nie dopisuje: i = σ(−3) ≈ 0,047, przy kandydacie g = tanh(1,1) ≈ 0,80. Nowa pamięć: c_t = 0,953 · 1 + 0,047 · 0,80 ≈ 0,991. Z bramką wyjścia o = σ(1) ≈ 0,731 stan ukryty to h_t = 0,731 · tanh(0,991) ≈ 0,554.
Teraz 50 takich kroków. Przy f ≈ 0,953 z pierwotnej informacji zostaje 0,953^50 ≈ 0,088, a gdy sieć podniesie f do σ(5) ≈ 0,993 — aż 0,993^50 ≈ 0,71. W prostej RNN z wzmocnieniem 0,5 na krok zostałoby 0,5^50 ≈ 9·10^(−16). Cena: w PyTorch nn.LSTM(100, 128) ma 4 · (128·100 + 128·128 + 2·128) = 117 760 parametrów, nn.GRU(100, 128) — 88 320, a zwykła nn.RNN o tych wymiarach — 29 440. LSTM to cztery „kopie” RNN (trzy bramki i kandydat), GRU — trzy.
W praktyce
- PyTorch:
nn.LSTM(input_size, hidden_size, num_layers=2, batch_first=True, dropout=0.2); wynik tooutput, (h_n, c_n).nn.GRUzwraca tylkooutput, h_n. - Typowe rozmiary stanu: 64–512; 1–3 warstwy. Więcej warstw rzadko pomaga bez dużej ilości danych.
- Obcinaj gradient (
clip_grad_norm_) — LSTM łagodzi zanikanie, ale nie chroni przed eksplozją. - Popularna sztuczka: inicjalizacja biasu bramki zapominania wartością dodatnią (np. 1), żeby sieć na starcie raczej pamiętała niż zapominała.
- GRU wybierz przy mniejszych danych lub ograniczonej pamięci, LSTM przy dłuższych zależnościach — ale ostatecznie zdecyduj na zbiorze walidacyjnym.
Najczęstsze pytania
- LSTM czy GRU — co wybrać?
- Porównania empiryczne nie wskazują jednoznacznego zwycięzcy. GRU jest mniejszy i szybszy, LSTM ma osobną pamięć i bywa lepszy przy bardzo długich zależnościach. W praktyce warto sprawdzić oba.
- Dlaczego LSTM nie zapomina tak szybko jak RNN?
- Bo stan komórki aktualizuje się przez dodawanie, a nie przez wielokrotne przepuszczanie przez nieliniowość. Gdy bramka zapominania jest bliska 1, informacja i gradient przechodzą przez kolejne kroki prawie bez zmian.
- Czy transformery całkowicie zastąpiły LSTM?
- W przetwarzaniu języka w dużej skali — tak. W szeregach czasowych, systemach wbudowanych i przetwarzaniu strumieniowym LSTM i GRU wciąż są używane, bo mają stały koszt na krok i małe wymagania pamięciowe.
Źródła
- Hochreiter, Schmidhuber „Long Short-Term Memory”, Neural Computation 9(8), 1997.
- Gers, Schmidhuber, Cummins „Learning to Forget: Continual Prediction with LSTM”, Neural Computation 12(10), 2000.
- Cho i in. „Learning Phrase Representations using RNN Encoder–Decoder for Statistical Machine Translation”, EMNLP 2014, arXiv:1406.1078.
- Goodfellow, Bengio, Courville „Deep Learning”, MIT Press, 2016, rozdz. 10.10.
- Zhang i in. „Dive into Deep Learning”, d2l.ai, rozdz. 10 („Modern Recurrent Neural Networks”).