ML Atlas

07 · Architektury · 5 min czytania · aktualizacja

RNN, LSTM czy transformer — którą architekturę wybrać do sekwencji?

W skrócie

Transformer najlepiej łączy odległe elementy sekwencji i uczy się równolegle. LSTM sprawdza się przy strumieniach i małej pamięci, zwykła RNN przy krótkich.

Co to jest

Do tekstu i innych długich sekwencji, gdy masz dużo danych lub gotowy model wstępnie wytrenowany, wybierz transformer; LSTM lub GRU — gdy przetwarzasz strumień krok po kroku przy stałej, małej pamięci (urządzenia, sygnały, sterowanie) albo masz mało danych i krótkie sekwencje; zwykłą sieć rekurencyjną (RNN) — praktycznie tylko jako punkt odniesienia lub do bardzo krótkich zależności. Kluczowa różnica to sposób przenoszenia informacji: RNN i LSTM przekazują ją przez kolejne kroki, transformer łączy każdy element z każdym bezpośrednio.

Zwykła RNN czyta sekwencję element po elemencie i za każdym razem aktualizuje jeden wektor stanu ukrytego: nowy stan = f(stary stan, bieżące wejście). Cała pamięć o przeszłości musi się zmieścić w tym wektorze i przetrwać wielokrotne przekształcenia.

LSTM dodaje do tego osobną „komórkę pamięci” i trzy bramki, które uczą się, co zapisać, co zapomnieć i co odczytać. Transformer w ogóle rezygnuje z rekurencji: samouwaga pozwala każdemu elementowi bezpośrednio odczytać informację z dowolnego innego.

Mechanizm — dlaczego tak działa

Zanikający gradient w RNN. Żeby nauczyć się zależności między krokiem 1 a krokiem T, gradient musi przejść wstecz przez T mnożeń przez macierz wag i pochodną funkcji aktywacji. Jeśli te czynniki są zwykle mniejsze od 1, gradient maleje wykładniczo; jeśli większe — eksploduje. Bengio, Simard i Frasconi (1994) pokazali, że to problem fundamentalny: zwykła RNN praktycznie nie uczy się zależności dłuższych niż kilkanaście–kilkadziesiąt kroków.

Jak LSTM to obchodzi. Komórka pamięci LSTM jest aktualizowana addytywnie: nowa komórka = bramka zapominania × stara komórka + bramka wejścia × nowa treść (Hochreiter i Schmidhuber, 1997; bramkę zapominania dodali Gers, Schmidhuber i Cummins, 2000). Gdy bramka zapominania jest bliska 1, informacja i gradient płyną przez wiele kroków prawie bez strat. To nie usuwa problemu całkowicie, ale wydłuża zasięg pamięci z kilkunastu do dziesiątek, a przy dobrym treningu nawet setek kroków. GRU to uproszczona wersja z dwiema bramkami i podobną skutecznością.

Dlaczego transformer wygrał w NLP. Samouwaga łączy dowolne dwa elementy sekwencji ścieżką długości 1, więc odległość nie utrudnia uczenia. Wszystkie pozycje liczy się równolegle, co na GPU oznacza dużo szybszy trening niż sekwencyjne przechodzenie RNN (Vaswani i in., 2017). Ceną jest koszt rosnący kwadratowo z długością sekwencji oraz brak wbudowanej kolejności — trzeba ją dostarczyć kodowaniem pozycji.

Gdzie rekurencja wciąż ma sens. RNN i LSTM przetwarzają strumień przy stałym koszcie pamięci na krok, niezależnie od długości historii, i naturalnie działają online. Transformer przy generowaniu musi pamiętać klucze i wartości wszystkich poprzednich tokenów. Stąd odrodzenie modeli rekurencyjnych i przestrzeni stanów dla bardzo długich sekwencji.

Na przykładzie

Dwa eksperymenty w PyTorch, trzy ziarna losowości, Adam (0,001), minipaczki po 32. Rekurencyjne sieci mają stan ukryty 64; transformer — dwa bloki, wymiar 64, cztery głowy, uczone osadzenia pozycji.

Najpierw cyfry Digits 8×8 (1347 treningowych, 450 testowych, random_state=0) podawane jako sekwencja: albo 8 wierszy po 8 pikseli, albo 64 pojedyncze piksele. 40 epok, trafność na teście:

SekwencjaRNNLSTMTransformer
8 kroków (wiersze)0,9690,9570,976
64 kroki (piksele)0,7860,8030,925
Parametry (64 kroki)493817 80271 818

Przy 8 krokach wszystkie trzy radzą sobie podobnie, a najprostsza RNN wypada nawet lepiej od LSTM. Przy 64 krokach rekurencyjne sieci tracą 15–18 punktów, bo muszą przenieść informację o pierwszych wierszach obrazka przez kilkadziesiąt kroków; transformer traci tylko 5 punktów. Uczciwie: transformer ma tu 4–15 razy więcej parametrów, a LSTM przy dłuższym treningu prawdopodobnie by się poprawił.

Drugi test mierzy samą pamięć: sekwencja T losowych bitów, a zadanie to podać pierwszy bit. 2000 sekwencji treningowych, 1000 testowych, 30 epok:

Długość TRNNLSTMTransformer
101,0001,0001,000
500,6671,0001,000
2000,6500,5231,000

Średnia RNN przy T = 50 i T = 200 ukrywa rozkład zero-jedynkowy: w dwóch próbach na trzy sieć zgaduje (ok. 0,5), w jednej rozwiązuje zadanie. To typowy obraz zanikającego gradientu — sukces zależy od szczęśliwej inicjalizacji. Niespodzianka przy T = 200: LSTM nie znalazł rozwiązania w żadnej z trzech prób (0,518–0,527). Bramki wydłużają pamięć, ale nie bez końca — przy tym budżecie treningu sygnał z pierwszego kroku ginie w 200 krokach szumu. Transformer rozwiązuje zadanie zawsze, bo uwaga może patrzeć wprost na pozycję numer 1, niezależnie od odległości.

Dane: Digits (ręcznie pisane cyfry 8×8)

W praktyce

Reguła wyboru:

  • Tekst, kod, długie zależności, dostępny model wstępnie wytrenowany → transformer: nn.TransformerEncoder(nn.TransformerEncoderLayer(d_model=128, nhead=4, batch_first=True), num_layers=2) albo gotowy model z Hugging Face.
  • Strumień danych na żywo, mała pamięć, krótkie lub średnie sekwencje, mało danych → nn.LSTM(input_size, hidden_size=64, batch_first=True) lub nn.GRU.
  • Zwykła nn.RNN → jako punkt odniesienia albo przy zależnościach na kilka kroków; przy dłuższych zawsze sprawdź LSTM/GRU.
  • RNN i LSTM: przycinaj gradient (torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)), żeby uniknąć eksplozji.
  • Transformer: pamiętaj o kodowaniu pozycji i masce przyczynowej przy przewidywaniu przyszłości (nn.Transformer.generate_square_subsequent_mask(T)); koszt pamięci rośnie z kwadratem długości.
  • Szeregi czasowe tabelaryczne: zanim sięgniesz po sieć, sprawdź boosting na cechach opóźnionych — często wygrywa przy mniejszym koszcie.

Najczęstsze pytania

Czy LSTM jest jeszcze używany?
Tak, choć rzadziej w NLP. Sprawdza się w przetwarzaniu sygnałów, rozpoznawaniu mowy na urządzeniach, sterowaniu i prognozowaniu, gdzie liczy się działanie strumieniowe i mały model. Przy dużych zbiorach tekstu wyparły go transformery.
Czym różni się GRU od LSTM?
GRU łączy komórkę pamięci ze stanem ukrytym i ma dwie bramki zamiast trzech, więc ma mniej parametrów i liczy się szybciej. W większości zadań obie dają podobne wyniki; GRU bywa lepsza przy mniejszych zbiorach, LSTM przy bardzo długich zależnościach.
Dlaczego transformer jest szybszy w treningu, skoro ma kwadratowy koszt?
Bo wszystkie pozycje liczy równolegle, a RNN musi czekać na wynik poprzedniego kroku. Dla typowych długości (setki–tysiące tokenów) równoległość na GPU wygrywa z kwadratowym kosztem. Przy bardzo długich sekwencjach kwadratowy koszt zaczyna dominować i stosuje się uwagę rzadką lub lokalną.

Źródła

  • Bengio Y., Simard P., Frasconi P. „Learning long-term dependencies with gradient descent is difficult”, IEEE Transactions on Neural Networks 5(2), 1994, s. 157–166.
  • Hochreiter S., Schmidhuber J. „Long Short-Term Memory”, Neural Computation 9(8), 1997, s. 1735–1780.
  • Gers F. A., Schmidhuber J., Cummins F. „Learning to Forget: Continual Prediction with LSTM”, Neural Computation 12(10), 2000, s. 2451–2471.
  • Cho K. i in. „Learning Phrase Representations using RNN Encoder–Decoder for Statistical Machine Translation”, EMNLP 2014.
  • Vaswani A. i in. „Attention Is All You Need”, NeurIPS 2017.
  • Goodfellow I., Bengio Y., Courville A. „Deep Learning”, MIT Press 2016, rozdz. 10.

Zobacz też