Нейронные сети Рекуррентные сети, LSTM и GRU
0%

Рекуррентные сети, LSTM и GRU

Рекуррентные сети, LSTM и GRU

В статье о свёрточных сетях мы разобрали архитектуру для данных с пространственной решёткой. Теперь — про данные, у которых есть порядок и переменная длина: текст, речь, логи кликов, показания датчиков, котировки, ДНК. Другой индуктивный сдвиг — другая архитектура. Даже если вы собираетесь работать только с трансформерами, этот материал нужен: внимание выросло ровно из попытки залатать узкое место seq2seq на LSTM, а современные state-space модели (Mamba, RWKV) — это рекуррентность, вернувшаяся с новой стороны.

1. Почему обычная сеть не справляется с последовательностью

Возьмём задачу «определить тональность отзыва». Отзывы разной длины: 5 слов и 500. MLP из статьи о перцептроне принимает вектор фиксированной размерности. Наивные обходы и их провалы: обрезать до 200 токенов и подать всё разом — параметр для слова на позиции 7 никак не связан с параметром для позиции 143, сеть должна выучить понятие «не» отдельно для каждой позиции, данных на это не хватит никогда; усреднить эмбеддинги — порядок теряется, «не хорошо, а плохо» и «не плохо, а хорошо» дают один вектор; 1D-свёртка — уже лучше, есть разделение весов вдоль времени, но рецептивное поле ограничено глубиной, а согласование через 40 слов требует нереальной глубины.

Рекуррентная сеть решает это иначе. Ключевая идея — состояние:

Читаем последовательность по одному элементу и поддерживаем вектор $h_t$ — сжатый конспект всего, что видели до момента $t$. На каждом шаге применяем одну и ту же функцию обновления: $h_t = f_\theta(h_{t-1}, x_t)$.

Это ровно то, что делает человек, слушая лекцию: он не хранит стенограмму, а держит в голове «состояние понимания» и обновляет его каждой новой фразой. Здесь два индуктивных сдвига: разделение параметров во времени (как в CNN — по пространству), дающее инвариантность к сдвигу и работу с любой длиной, и марковское сжатие — всё прошлое обязано поместиться в $h_t$ фиксированной размерности. Второе — и сила (константная память), и главное ограничение; вернёмся к нему в разделе о seq2seq.

2. Vanilla RNN: определение и развёртка

Классическая формулировка (Elman, 1990):

$$h_t = \tanh(W_{hh} h_{t-1} + W_{xh} x_t + b_h), \qquad \hat{y}_ t = W_{hy} h_t + b_y$$

где $h_t \in \mathbb{R}^{d}$, $x_t \in \mathbb{R}^{m}$, $h_0$ обычно нулевой. Параметров в рекуррентной части $d^2 + dm + d$, и они не зависят от длины последовательности. Рекурсия — цикл в графе вычислений; чтобы применить обратное распространение, цикл разворачивают: сеть длины $T$ превращается в обычную feed-forward сеть глубины $T$ со связанными весами.

Развёртка RNN во времени и затухание градиента

Именно эта картинка объясняет всё дальнейшее: RNN на последовательности из 200 токенов — это сеть глубиной 200 слоёв, где все слои используют одну и ту же матрицу. Форвард-проход на голом NumPy, чтобы не осталось магии:

import numpy as np

def rnn_forward(xs, h0, W_xh, W_hh, b_h, W_hy, b_y):
    """xs: (T, m). Возвращает состояния hs и логиты."""
    hs = np.zeros((len(xs) + 1, h0.shape[0]))
    hs[-1] = h0                      # трюк: индекс -1 хранит h_0
    logits = []
    for t in range(len(xs)):
        # одна и та же W_hh на каждом шаге — в этом весь смысл
        hs[t] = np.tanh(W_xh @ xs[t] + W_hh @ hs[t - 1] + b_h)
        logits.append(W_hy @ hs[t] + b_y)
    return hs, np.array(logits)

Сложность. Один шаг — $O(d^2 + dm)$ операций; вся последовательность — $O(T(d^2 + dm))$ по времени и $O(T d)$ по памяти на батч-элемент (все $h_t$ нужны для обратного прохода). Критично другое: последовательная глубина $O(T)$ — шаг $t$ нельзя посчитать, не досчитав $t-1$, и GPU простаивает. Это тот самый структурный недостаток, который убьёт RNN, когда появятся трансформеры с их параллельным по времени вниманием. Форматы задач, которые на RNN решают:

3. BPTT: обратное распространение сквозь время

Развернули — значит, работает обычный backprop (см. обратное распространение), с одной особенностью: градиенты по общим весам суммируются по всем шагам.

$$\frac{\partial L}{\partial W_{hh}} = \sum_{t=1}^{T} \frac{\partial L_t}{\partial W_{hh}}, \qquad \frac{\partial L_T}{\partial h_t} = \frac{\partial L_T}{\partial h_T} \prod_{k=t+1}^{T} \frac{\partial h_k}{\partial h_{k-1}}$$

Ключевой множитель — якобиан одного шага $J_k = \partial h_k / \partial h_{k-1} = \mathrm{diag}(1 - h_k^2), W_{hh}^{\top}$. Реализация BPTT — почти дословный перевод формул:

def rnn_bptt(xs, targets, hs, probs, W_hh, W_xh, W_hy):
    """probs: (T, C) softmax-вероятности; targets: (T,) индексы классов."""
    dW_xh, dW_hh, dW_hy = (np.zeros_like(W_xh), np.zeros_like(W_hh),
                           np.zeros_like(W_hy))
    db_h, db_y = np.zeros(W_hh.shape[0]), np.zeros(W_hy.shape[0])
    dh_next = np.zeros(W_hh.shape[0])
    for t in reversed(range(len(xs))):
        dy = probs[t].copy()
        dy[targets[t]] -= 1.0             # grad кросс-энтропии по логитам
        dW_hy += np.outer(dy, hs[t]); db_y += dy
        dh = W_hy.T @ dy + dh_next        # вклад «своего» выхода + из будущего
        dz = (1 - hs[t] ** 2) * dh        # прошли назад через tanh
        dW_xh += np.outer(dz, xs[t])
        dW_hh += np.outer(dz, hs[t - 1])  # веса общие → аккумулируем
        db_h  += dz
        dh_next = W_hh.T @ dz             # передаём градиент в прошлое
    return dW_xh, dW_hh, dW_hy, db_h, db_y

Truncated BPTT. Для потоков в миллионы шагов (языковая модель на корпусе, телеметрия) полная развёртка невозможна: память $O(T)$. Разворачивают окно $k$ шагов (обычно 32–256), делают шаг оптимизатора, переносят $h$ в следующее окно как константу (h.detach()) и обрезают там граф. Обучение видит зависимости только внутри окна, зато forward-состояние едет дальше. Забыть detach() — классический баг: граф растёт бесконечно, на третьей итерации OOM.

4. Почему vanilla RNN не учится: строгий разбор

Возьмём норму произведения якобианов. Так как $|\tanh’| \le 1$:

$$\left| \prod_{k=t+1}^{T} J_k \right| \le \left(\sigma_{\max}(W_{hh})\right)^{T-t}$$

Отсюда результат Pascanu, Mikolov, Bengio (arXiv:1211.5063): при $\sigma_{\max}(W_{hh}) < 1$ градиент гарантированно затухает экспоненциально по дистанции, а при наибольшем собственном значении $> 1$ возможен экспоненциальный взрыв. Экспонента с показателем 100 не прощает ничего: при коэффициенте 0.9 за шаг через 100 шагов остаётся $2.6 \times 10^{-5}$ сигнала. Проверим численно:

rng, d = np.random.default_rng(0), 100
for scale in (0.8, 1.3):
    W = rng.normal(0, scale / np.sqrt(d), (d, d))   # спектральный радиус ≈ scale
    v = rng.normal(0, 1, d)
    for _ in range(100):
        v = W.T @ (v * 0.9)                         # 0.9 — типичное tanh'
    print(scale, f"{np.linalg.norm(v):.1e}")        # 0.8 → 1e-22, 1.3 → 1e+14

Практические следствия. Взрыв лечится симптоматически — обрезкой нормы градиента: направление сохраняется, ограничивается длина шага (torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0), типичные значения 0.25–5). Затухание симптоматически не лечится: сигнал не «уменьшен» — он смешан с шумом других шагов и потерян, надо менять саму форму рекуррентности так, чтобы существовал путь с якобианом, близким к единичному. Полумеры, которые пробовали: ортогональная инициализация $W_{hh}$ (Saxe et al., arXiv:1312.6120) — спектральный радиус ровно 1 на старте; IRNN — ReLU плюс $W_{hh}=I$ (Le et al., arXiv:1504.00941). Работает, но хрупко. Победила архитектурная идея.

5. LSTM: конвейерная лента и вентили

Hochreiter и Schmidhuber предложили LSTM в 1997 году (PDF оригинала) — за 20 лет до трансформеров. Интуиция: разделим «память» и «рабочее представление». Заведём вектор состояния ячейки $c_t$ — конвейерную ленту, которая едет сквозь время и по умолчанию ничего с собой не делает. Вместо полного пересчёта состояния каждый шаг (как в vanilla RNN, где $h_t$ целиком переписывается через tanh) сеть учится трём решениям: что из ленты стереть (forget gate), что на неё дописать (input gate плюс кандидат), что прочитать наружу (output gate). Вентиль — вектор из $[0,1]^d$ от сигмоиды, поэлементно умножаемый на поток: 0 — закрыт, 1 — открыт. Он обучаемый и зависит от входа, то есть решение «забыть/запомнить» принимается контекстно, а не раз и навсегда.

$$ \begin{aligned} f_t &= \sigma(W_f [h_{t-1}, x_t] + b_f) &\quad& \text{сколько оставить от } c_{t-1} \ i_t &= \sigma(W_i [h_{t-1}, x_t] + b_i) &\quad& \text{сколько записать нового} \ \tilde{g}_ t &= \tanh(W_g [h_{t-1}, x_t] + b_g) &\quad& \text{что именно записать} \ o_t &= \sigma(W_o [h_{t-1}, x_t] + b_o) &\quad& \text{что показать наружу} \ c_t &= f_t \odot c_{t-1} + i_t \odot \tilde{g}_ t &\quad& \textbf{обновление ленты} \ h_t &= o_t \odot \tanh(c_t) &\quad& \text{видимое состояние} \end{aligned} $$

Схема ячейки LSTM

Почему это чинит градиент

Смотрим на путь по ленте, игнорируя ветви через вентили: $\partial c_t / \partial c_{t-1} = \mathrm{diag}(f_t)$. Ни матрицы весов, ни производной tanh в этом пути нет — только поэлементное умножение на $f_t$. Если сеть выучила $f_t \approx 1$ для нужных координат, градиент за 500 шагов умножается на $\approx 1$ и доходит неповреждённым: это авторский constant error carousel. Заметьте структурное родство: $c_t = f_t \odot c_{t-1} + \ldots$ — аддитивное обновление, ровно как residual-связь $x + F(x)$ в ResNet; одна идея «сделай тождественное отображение путём по умолчанию» решает проблему глубины и во времени, и в пространстве. Разные режимы вентилей дают качественно разное поведение памяти:

Инициализация forget gate — самое дешёвое улучшение

При стандартной инициализации $b_f = 0$ вентиль стартует с $f_t \approx 0.5$: за 20 шагов остаётся $10^{-6}$ сигнала, то есть сеть начинает с короткой памяти, а долгой ей надо научиться — градиентом, которого нет. Лечится одной строкой: $b_f = 1$ (или 2), тогда $f_t \approx 0.73$–$0.88$ на старте — по Jozefowicz et al. едва ли не самое влиятельное решение во всей настройке LSTM (PMLR v37).

import torch, torch.nn as nn

lstm = nn.LSTM(128, 256, num_layers=2, batch_first=True)
for name, p in lstm.named_parameters():
    if "bias" in name:
        p.data.fill_(0.0)
        p.data[256:512].fill_(1.0)       # порядок гейтов в PyTorch: i, f, g, o
    elif "weight_hh" in name:            # ортогональная W_hh поблочно на гейт
        for k in range(4):
            nn.init.orthogonal_(p.data[k * 256:(k + 1) * 256])

6. GRU: то же самое, но дешевле

Cho et al., 2014 (arXiv:1406.1078) убрали отдельное состояние ячейки и свели четыре блока к трём:

$$ \begin{aligned} z_t &= \sigma(W_z [h_{t-1}, x_t]) &\quad& \text{update gate} \ r_t &= \sigma(W_r [h_{t-1}, x_t]) &\quad& \text{reset gate} \ \tilde{h}_ t &= \tanh(W_h [r_t \odot h_{t-1}, x_t]) \ h_t &= (1 - z_t) \odot h_{t-1} + z_t \odot \tilde{h}_ t \end{aligned} $$

Три отличия от LSTM: одно состояние вместо двух (нет output gate, состояние всегда полностью экспонируется наружу); связанные вентили — вместо независимых $f$ и $i$ один $z$ с ограничением «сколько забыл, столько и записал»; reset gate до нелинейности, позволяющий начать с чистого листа на новом сегменте. Параметров при $d$ скрытых и $m$ входных: LSTM — $4(d^2 + dm)$, GRU — $3(d^2 + dm)$, то есть GRU на 25 % легче и быстрее.

Что выбирать? Chung et al. (arXiv:1412.3555) и ablation Greff et al. LSTM: A Search Space Odyssey (arXiv:1503.04069, 5400 запусков) дают одинаковый ответ: статистически значимой разницы в качестве нет, и ни один из популярных вариантов LSTM (peephole, объединённые вентили) не превосходит базовый. Правило: ограниченный бюджет и on-device → GRU; очень длинные зависимости и большой корпус → LSTM; в любом случае forget bias и clipping важнее выбора между ними. Нюанс PyTorch: в nn.GRU reset gate применяется после матричного умножения ($n_t = \tanh(W_{in}x_t + r_t \odot (W_{hn}h_{t-1} + b_{hn}))$), а не до, как в статье — ради одного слитного GEMM. Переносите веса из чужой реализации — проверьте эту деталь, иначе получите тихо неверную модель.

7. Реализация: от ручной ячейки до продакшн-цикла

Сначала ячейка руками — она должна совпадать с эталонной до 1e-6:

import torch

def lstm_cell(x, state, W_ih, W_hh, b_ih, b_hh):
    """x: (B, m); state: (h, c), каждый (B, d). Один шаг LSTM."""
    h_prev, c_prev = state
    # ОДИН матмул на все четыре гейта — так делают все быстрые реализации
    gates = x @ W_ih.T + b_ih + h_prev @ W_hh.T + b_hh      # (B, 4d)
    i, f, g, o = gates.chunk(4, dim=1)                      # порядок PyTorch
    i, f, o = torch.sigmoid(i), torch.sigmoid(f), torch.sigmoid(o)
    c = f * c_prev + i * torch.tanh(g)
    return o * torch.tanh(c), c

ref = torch.nn.LSTMCell(8, 16)                              # сверка с эталоном
x, h0, c0 = torch.randn(4, 8), torch.randn(4, 16), torch.randn(4, 16)
mine = lstm_cell(x, (h0, c0), ref.weight_ih, ref.weight_hh, ref.bias_ih, ref.bias_hh)
assert torch.allclose(mine[0], ref(x, (h0, c0))[0], atol=1e-6)

Теперь боевой вариант — классификатор с паддингом и упаковкой:

import torch.nn as nn
from torch.nn.utils.rnn import pack_padded_sequence, pad_packed_sequence

class SeqClassifier(nn.Module):
    def __init__(self, vocab, emb=128, hid=256, n_layers=2, n_cls=2, p=0.3):
        super().__init__()
        self.emb = nn.Embedding(vocab, emb, padding_idx=0)
        self.rnn = nn.LSTM(emb, hid, n_layers, batch_first=True,
                           bidirectional=True, dropout=p)   # dropout МЕЖДУ слоями
        self.drop = nn.Dropout(p)
        self.head = nn.Linear(hid * 2, n_cls)               # *2 из-за bidirectional

    def forward(self, tokens, lengths):
        e = self.drop(self.emb(tokens))                     # (B, T, emb)
        # упаковка: cuDNN не считает шаги по паддингу, а состояние
        # не «размывается» нулями в конце коротких примеров
        packed = pack_padded_sequence(e, lengths.cpu(),
                                      batch_first=True, enforce_sorted=False)
        out, (h_n, _) = self.rnn(packed)
        out, _ = pad_packed_sequence(out, batch_first=True)  # (B, T, 2*hid)
        # h_n: (n_layers*2, B, hid) — берём оба направления последнего слоя
        last = torch.cat([h_n[-2], h_n[-1]], dim=1)          # (B, 2*hid)
        return self.head(self.drop(last))

def train_step(model, batch, opt, loss_fn):
    tokens, lengths, y = batch
    opt.zero_grad(set_to_none=True)
    loss = loss_fn(model(tokens, lengths), y)
    loss.backward()
    nn.utils.clip_grad_norm_(model.parameters(), 1.0)  # без этого будет NaN
    opt.step()
    return loss.item()

Про производительность. nn.LSTM на CUDA уходит в слитное ядро cuDNN: все гейты — один GEMM, элементные операции сплавлены. Свой Python-цикл по LSTMCell медленнее в 10–50 раз: на каждый шаг — запуск десятка мелких ядер. Нужен кастомный шаг (скажем, своё внимание внутри рекуррентности) — пишите под torch.jit.script или torch.compile. Быстрый путь cuDNN требует contiguous-тензоров и dropout только между слоями; проверять — профайлером, а не на глаз.

8. Seq2seq и то самое узкое место

Sutskever et al., 2014 (arXiv:1409.3215): энкодер-LSTM сжимает предложение в один вектор, декодер-LSTM разворачивает его в перевод.

Две проблемы видны прямо на диаграмме. Информационное бутылочное горлышко: качество перевода падало с ростом длины предложения — фиксированный вектор физически не вмещает длинный вход. Ответ — внимание Bahdanau et al. (arXiv:1409.0473): декодер на каждом шаге смотрит на все состояния энкодера со взвешиванием. Отсюда прямая дорога к трансформерам, которые просто выбросили рекуррентность. Exposure bias: на обучении декодер видит правильную историю, на инференсе — свою собственную, с ошибками; смягчается scheduled sampling (arXiv:1506.03099), но полностью не лечится.

9. Регуляризация и нормализация в рекуррентных сетях

Приёмы из статьи об обучении сетей здесь работают не «как есть»:

  • Обычный dropout на рекуррентную связь вредит: независимая маска на каждом шаге превращает память в шум. Правильный вариант — variational dropout, одна маска на все шаги (Gal & Ghahramani, arXiv:1512.05287). Параметр dropout в nn.LSTM применяется только между слоями, не внутри времени.
  • BatchNorm плохо дружит с рекуррентностью: статистики зависят от шага $t$, а длины в батче разные. Стандарт — LayerNorm (по признакам внутри примера).
  • DropConnect на $W_{hh}$ + усреднённый SGD — рецепт AWD-LSTM (Merity et al., arXiv:1708.02182), несколько лет державший SOTA в языковом моделировании: образец того, как регуляризация бьёт размер модели. Родственный zoneout вместо зануления копирует предыдущее состояние — это дружественно к градиентному потоку.

10. Trade-offs: где рекуррентность всё ещё выигрывает

Свойство RNN / LSTM / GRU Трансформер
Время на последовательность $O(T d^2)$ — линейно $O(T^2 d + T d^2)$
Последовательная глубина $O(T)$ — плохо для GPU $O(1)$ — параллелится
Память на инференсе $O(d)$ — константа $O(T d)$ — KV-кэш растёт
Путь между токенами $i$ и $j$ $O(|i-j|)$ шагов $O(1)$
Стриминг с малой задержкой нативный требует чанкования

Читается так: трансформер выиграл тренировку, но рекуррентность не проиграла инференс. Отсюда — где RNN живы в проде:

  • Стриминговое распознавание речи. RNN-Transducer (Graves, arXiv:1211.3711) — основа онлайн-ASR на устройствах: выдаёт гипотезу по мере поступления звука, память константная.
  • On-device и edge. Keyword spotting, детекция голосовой активности, носимые устройства: GRU на 64 нейрона укладывается в микроконтроллер, где трансформер не поместится в принципе. Сюда же — короткие последовательности с жёстким latency SLA (сессия из 20 кликов, антифрод по транзакциям).
  • Возвращение рекуррентности. RWKV (arXiv:2305.13048), Mamba/SSM (arXiv:2312.00752), xLSTM (arXiv:2405.04517) — по сути LSTM, переписанная так, чтобы обучение тоже параллелилось (ассоциативный скан), а инференс остался рекуррентным. Понимание LSTM — прямая предпосылка к ним.

11. Типичные ошибки

  1. Забыть clip_grad_norm_. Обучение идёт, а на 3000-м шаге лосс становится NaN. Диагностика: логируйте total_norm, который возвращает функция обрезки.
  2. Забыть detach() состояния между окнами TBPTT. Граф растёт, память кончается.
  3. Паддинг без упаковки или маски. Модель усредняет нули — качество тихо деградирует на 5–15 %, никакой ошибки не возникает. Сюда же — взять output[:, -1, :] вместо h_n: для коротких примеров это состояние после паддинга.
  4. Bidirectional-энкодер в causal-задаче. Двунаправленная LSTM видит будущее — в языковом моделировании или онлайн-предсказании это утечка целевой переменной: валидация блестящая, прод не работает. Сюда же — mean/std временного ряда, посчитанные вместе с тестом.
  5. Dropout внутри рекуррентной связи со свежей маской каждый шаг. Не сходится.
  6. Ожидать от LSTM памяти на 10 000 шагов. Горизонт — сотни шагов, в лучшем случае пара тысяч; нужно больше — нужно внимание или иерархия.
  7. Сортировать батч по длине и забыть восстановить порядок меток. enforce_sorted=False снимает вопрос — используйте его.

12. Мини-итог

  • RNN — это weight sharing вдоль оси времени плюс скрытое состояние как сжатый конспект прошлого; развёртка превращает её в очень глубокую сеть.
  • BPTT упирается в произведение $T$ якобианов: экспоненциальное затухание или взрыв — не баг настройки, а математика. Взрыв лечится обрезкой градиента, затухание — только архитектурой: нужен путь с якобианом $\approx I$.
  • LSTM делает такой путь явно: $c_t = f_t \odot c_{t-1} + i_t \odot \tilde{g}_ t$, где $\partial c_t / \partial c_{t-1} = \mathrm{diag}(f_t)$ — тот же принцип, что и residual-связь, только во времени. GRU — та же идея на 25 % дешевле, и выбор между ними менее важен, чем clipping, forget bias и работа с паддингом.
  • Рекуррентность проиграла трансформерам по параллелизму обучения, но выигрывает по инференсу: линейное время, константная память, нативный стриминг.

Источники

Что дальше

Рекуррентность упирается в два ограничения: последовательное вычисление и сжатие всего контекста в один вектор. Механизм внимания снимает второе, а отказ от рекуррентности — первое. Дальше — Внимание и трансформеры: self-attention с нуля, multi-head, позиционные кодировки и то, почему именно эта архитектура масштабируется до больших языковых моделей.

Нашли неточность? Выделите фрагмент текста — рядом появится жучок.

Нужен разбор именно вашей ситуации?

Статья описывает общий случай. Если у вас частный — можно разобрать его отдельно, платно. А если не хватает целого материала, предложите тему: её оплачивают вскладчину, и она выходит открытой для всех.

Доска запросов