Всем привет!

Не давно сделал свою Gated RNN, то есть с другой математикой, ни как у LSTM, GRU и подобного.

Я хочу (для вас) разобрать её теоретически, практически, замерить (я не умею замерять так что буду замерять как могу), ну и конечно же расскажу плюсы и минусы.

То, что будет в статье

Э‑э-э, плохое название для заголовка, но вот те заголовки которые вы сейчас будете встречать (в правильном порядке):

  1. Теория.

  2. Практика (без замеров).

  3. Бенчмаркинг.

  4. Плюсы и минусы моей сети.

  5. Вывод.

Это моя первая статья похожая на реально научную, так что могут быть недочёты.

Теория

Начнём с теории и формул.

Я сделал несколько нестандартных решений:

  1. softsignи его моя версия за место tanh и sigmoid.

  2. x_{t}+ h_{t-1} за место конкатенирования.

На самом деле их много чем два, но перейдем к теории.

И так, для справки напишу формулу softsign:

softsign(x) = \frac{x}{1 + |x|}

Всё просто: x делим на его модуль + 1.

Теперь я хочу показать мою формулу softsign:

softsign_{scaled}(x) = \frac{1 + softsign(x)}{2}

Эта формула мне нужна для замены sigmoid (сигмоида выдает диапазон от 0 до 1, а обычный софтсайн — от -1 до 1, а мне нужно было от 0 до 1).

Показываю первую формулу для своеобразного «насыщения» (нужно для более лучшего обобщения) x_t:

x_{t,new} = softsign(x_{t,old} + h_{t-1}) \alpha

То есть, x_t теперь это x_t, если что. Работает просто — x_t (старый) суммируем с h_{t-1} и сумму пропускаем через softsign умножаем на обучаемый параметр \alpha (я его на 2 с начало ставлю, вроде так лучше по качеству и обобщению).

И так, показываю формулу для гейта forget (f_t):

f_t = softsign_{scaled}(W_f (x_t + h_{t-1}) + b_f)

В общем, это как из обычного LSTM, но без конкатенирование (заменил на +) и с моим softsign_{scaled}.

Теперь нам нужен гейт input (i_t):

i_t = softsign_{scaled}((W_{ix}x_t + b_{ix}) + (W_{ih}h_{t-1} + b_{ih}))

Работает так:

Прогоняем x_t и h_{t-1} через два разных линейных слоёв со смещением, суммируем оба результата и прогоняем сумму через softsign_{scaled}.

Сейчас я запишу вычисление кандидата:

\tilde{h}_t = softsign_{scaled}(f_t \odot h_{t-1} + i_t \odot x_t)

То есть, просто суммирование поэлементного f_t на h_{t-1} и i_t на x_t, а потом прогоняем через scaled softsign.

Я сам сомневаюсь в таком вычисление, но пока что это самый рабочий вариант (для моих задач).

Потом вычисляем output gate:

o_t = softsign(W_o(x_t + h_{t-1}) + b_o)

У вас наверняка вопроса:

Почему тут не softsign_{scaled}?

Ну, я уже пробовал сделать наоборот — в вычислении кандидата обычный софтсайн, а в вычисление output гейта — скейлед софтсайн, но качество проседало аж до 23%.

Новый h_t вычисляем просто:

h_t = o_t \odot \tilde{h}_t

Скорее всего ещё один вопрос у читателей — где C_t, где CEC?

Ответ прост — я решил убрать эту всю мишуру (ладно, это не мишура) ещё на старте, и оно заработало, я подумал — «ну ладно, если работает — в принципе, пока не надо» и так и осталось по сей день (уже нет).

Это первый этап вычисления в моей Gated RNN.

Возможно вы спросите — а где же долгосрочная (Long‑Term) память?

Ну, вот щас покажу.

На выходе первого этапа идет:

H = (h_0, h_1, ..., h_L)

L тут это длина всей последовательности X.

Потом идет второй этап (одна формула, да):

H_{long} = \frac{H \cdot (H^T \cdot H)}{\sqrt{d_h}}

d_h— это размерность скрытого состояния.

Работает так:

Умножаем H^T на H — получаем что‑то вроде «матрицы внимания» (термин не к месту наверно, да?).

Умножаем H на эту самую «матрицу внимания» чтобы сделать размерность правильной и делим на \sqrt{d_h} чтобы не взорвать градиенты.

Всё, это вся долгосрочная память.

Если моя Gated RNN — это последний слой всей сети (ну или там дальше идет LayerNorm или классификатор) — то мы выдаём такой output:

Output = softsign(\sum_{i=1}^{L} H_{long,i})

Если что, \sum тут считает сумму каждой строки матрицы H_long и все результаты в один список. То есть возьмём пример: [[1, 2], [2, 3]]. \sum тут посчитает и выдаст такой результат: [3, 5]. То есть, 1 + 2 = 3, 2 + 3 = 5, собираем в список — готово.

Если же дальше идёт какой то слой — просто передаем H_{long} как есть, хотя можно и прогнать через softsign если надо.

Дальше в моей Gated RNN после этих двух этапов идёт LayerNorm.

Почему не RMSNorm и не BatchNorm?

С ними у меня качество не поднималось никуда, а даже опускалось (да!). А дальше может идти классификатор, но я решил не ставить потому что и так все работало я боялся переобучения или что‑то вроде того.

Это кажется, вся структура моей сети.

Если что — первый этап назвал SWM — Short Working Memory, а второй — LWM — Long Working Memory (я так назвал потому что не мог другое придумать на самом деле), в общем эта махина называется LSWM — Long‑Short Working Memory.

Теперь общая цепочка которую я написал у себя в коде:

\text{Embedding -> Short Working Memory -> Long Working Memory -> LayerNorm}

Конец теории! Время практики...

Практика (Без замеров)

Перейдем к практике!

Я решил сразу сделать достаточно сложную задачу — называют её «Multi‑hop branching».

Обычный multi‑hop — это «a = b = c, что такое a?» и модель в теории должна выдать «c», но как оказалось, для моей сети это была простая задача.

А branching multi‑hop — это типа «a = b, а ещё a = c. Какой a в начале был задан, а какой в конце?».

В общем, у меня было два инференса после обучения:

  1. Просто тест («a = b, a = c»), без всяких изменений.

  2. Тест, но на (как это пафосно называют) экстраполяцию длины — типа длину теста делают больше. Так что тут уже было вот так: «a = b = c, a = c = b».

На втором тесте моя сеть и всегда валила.

Изначально мне казалось что это проблема в слое LWM.

Пытался «решить» я так — с начало попытался за место деления на \sqrt{d_h} поставить LayerNorm (качество было больше, но всё равно валила), потом вообще решил H с начало пропускать через три матрицы — Q, K, V — без изменений.

В общем перепробовал я все адекватные на мой взгляд варианты, и я понял что LWM мне не чем не поможет.

Тогда я подумал‑подумал — и понял — я забыл поставить output gate (ну да...).

В общем спустя час ковыряний с output gate (то превращал в обычный линейный слой, то ставил \tilde{h}_t за место нормального x_t + h_t) я пришёл к выводу каким надо сделать output gate. Ну, в разделе «Теория.» к этому варианту и пришёл.

И вот резко моя сеть стала проходить эти multi‑hopы.

Код теста
import torch
import torch.nn as nn
import torch.optim as optim
import numpy as np

device = torch.device("cuda")

class LSWM(nn.Module):
    def __init__(self, vocab_size, d_model):
        super().__init__()

        self.d = d_model
        self.sd = d_model ** 0.5

        self.vocab_size = vocab_size
        self.embedding = nn.Embedding(vocab_size, d_model).to(device)

        self.W_f = nn.Linear(d_model, d_model).to(device)
        self.W_ix = nn.Linear(d_model, d_model).to(device)
        self.W_ih = nn.Linear(d_model, d_model).to(device)
        self.W_o = nn.Linear(d_model, d_model).to(device)
        self.a = nn.Parameter(torch.scalar_tensor(2)).to(device)

        self.norm = nn.LayerNorm(d_model).to(device)

    def softsign(self, x):
        return x / (1.0 + torch.abs(x))

    def softsign_scaled(self, x):
        return (1.0 + self.softsign(x)) / 2.0

    def forward(self, token_seq):
        batch_size, seq_len = token_seq.size()

        x_seq = self.embedding(token_seq)

        h_t = torch.zeros(batch_size, self.d).to(device)
        h = []

        for t in range(seq_len):
            x_t = x_seq[:, t, :]

            x_normed = (self.softsign(x_t + h_t)) * self.a

            f_t = self.softsign_scaled(self.W_f(x_normed + h_t))
            i_t = self.softsign_scaled(self.W_ix(x_t) + self.W_ih(h_t))
            o_t = self.softsign(self.W_o(x_normed + h_t))

            c = self.softsign_scaled(f_t * h_t + i_t * x_normed)

            h_t = o_t * c

            h.append(h_t)

        h = torch.stack(h, dim=1)
        res = torch.bmm(h, torch.bmm(h.transpose(-2, -1), h))
        h = res / self.sd

        res = self.softsign(h.sum(dim=1))

        return self.norm(res)

VOCAB_SIZE = 21
TOKEN_ARROW = 15
TOKEN_Q_1 = 16
TOKEN_Q_2 = 17

def generate_branching_batch(batch_size, epoch):
    x = np.zeros((batch_size, 7), dtype=np.int64)
    y = np.zeros(batch_size, dtype=np.int64)

    for i in range(batch_size):
        a, b, c = np.random.choice(15, 3, replace=False)

        ask_live = epoch % 2 == 0

        if ask_live:
            x[i] = [a, TOKEN_ARROW, b, a, TOKEN_ARROW, c,  TOKEN_Q_1]
            y[i] = b
        else:
            x[i] = [a, TOKEN_ARROW, b, a, TOKEN_ARROW, c,  TOKEN_Q_2]
            y[i] = c

    return torch.tensor(x).to(device), torch.tensor(y).to(device)

D_MODEL = 128
model = LSWM(vocab_size=VOCAB_SIZE, d_model=D_MODEL).to(device)
criterion = nn.CrossEntropyLoss().to(device)
optimizer = optim.Adam(model.parameters(), lr=0.0002, weight_decay=0.0099999)

acc = 0
epoch = 1
while epoch != 1001:
    inputs, targets = generate_branching_batch(64, epoch)

    optimizer.zero_grad()
    outputs = model(inputs)
    loss = criterion(outputs, targets)
    loss.backward()
    optimizer.step()

    if epoch % 500 == 0:
        preds = torch.argmax(outputs, dim=1)
        acc = (preds == targets).float().mean().item() * 100
        print(f"Loss: {loss.item():.4f} | Accuracy: {acc:.1f}%")

    epoch += 1

a, b, c = 3, 7, 12

model.eval()

with torch.no_grad():
    test_live = torch.tensor([[a, TOKEN_ARROW, b, a, TOKEN_ARROW, c, TOKEN_Q_1]]).to(device)
    pred_live = torch.argmax(model(test_live), dim=1).item()

    test_work = torch.tensor([[a, TOKEN_ARROW, b, a, TOKEN_ARROW, c, TOKEN_Q_2]]).to(device)
    pred_work = torch.argmax(model(test_work), dim=1).item()

    print("test 1:")
    print(f"a 1: {pred_live}")
    print(f"a 2: {pred_work}")

with torch.no_grad():
    test_live = torch.tensor([[b, TOKEN_ARROW, c, TOKEN_ARROW, a, b, TOKEN_ARROW, a, TOKEN_ARROW, c, TOKEN_Q_1]]).to(device)
    pred_live = torch.argmax(model(test_live), dim=1).item()

    test_work = torch.tensor([[b, TOKEN_ARROW, c, TOKEN_ARROW, a, b, TOKEN_ARROW, a, TOKEN_ARROW, c, TOKEN_Q_2]]).to(device)
    pred_work = torch.argmax(model(test_work), dim=1).item()

    print("test 2:")
    print(f"b 1: {pred_live}")
    print(f"b 2: {pred_work}")

Запускаю и...

Loss: 0.5694 | Accuracy: 89.1%
Loss: 0.1655 | Accuracy: 100.0%
test 1:
chain: 3, arrow, 7, and, 3, arrow, 12
need: 7 a 1: 7
need: 12 a 2: 12
test 2:
chain: 7, arrow, 12, arrow, 3, and, 7, arrow, 3, arrow, 12
need: 3 b 1: 3
need: 12 b 2: 12

Всё как и надо.

Если что «need: число» и «chain: цепочка» — это я уже к результату приписал чтобы было понятнее.

Кстати — я ещё попробовал в тест 2 добавлять «мусорные» токены (их мало было — всего 3, но даже 3 я считаю уже значительным изменением) — так же работало, на мое удивление.

Другие тесты опубликовывать не буду (на Гитхаб опубликую уже), но вот таблица:

Тест

Правильно?

Эпох

Multi‑hop braching

Да

1000

Multi‑hop (обычный)

Да

1000

Простая синусоида

Близко (надо 0.2440, а сеть выдала 0.2698)

600

Как видим, сеть достаточно правильно отвечает!

Я правда ещё не делал тесты на генерацию текста, но я буду обязан их сделать в обозримом будущем.

Бенчмаркинг

Время перейти к бенчмаркам!

Записывать буду в таблицу все результаты.

Правда, вот появилась проблема — на моём Google Colab я исчерпал лимиты на GPU, так что замерять буду на CPU.

Я решил замерять на том же multi‑hop branching тесте (первом где a = b, a = c) который у меня описан ранее в разделе «Практика (без замеров).».

Метрика

LSTM

GRU

LSWM (torch.compile)

Лосс в конце обучения (400 эпох).

0.8333

0.2362

0.1666

Качество в конце обучения (400 эпох).

71.9%

96.9%

100.0%

Результат сети (1).

7

7

7

Результат сети (2).

12

12

12

Количество параметров.

68 609

68 880

68 609

Скорость обучения (400 эпох).

4 сек

4 сек

8 сек

Weight decay

0.0099999

0.0099999

0.0099999

Learning rate

0.0004

0.0004

0.0004

Hidden size

91

105

128

Как видим, LSWM обходит всех по точности, а GRU и LSTM — по скорости обучения, но их объединяет одно — у них всех ответы одинаково правильные.

Я не хочу делать второй бенчмарк на второй тест (там практически всё так же по скорости и всему остальному), так что дам результаты:

LSTM

GRU

LSWM

3

3

3

12

12

12

В общем это подтверждает то что они при любом случае выдадут одинаковые ответы после обучения на этой задаче.

Я решил изменить первую цепочку второго теста на такую цепочку:
«b — c — a — b — a».

Протестировал и я понял — моя сеть чувствительна к сиду (рандома).

Тогда я решил найти оптимальный вариант математики моей LSWM чтобы убрать чувствительность к сиду (рандома).

В итоге я сделал изменения:

\tilde{c}_t = softsign((W_cx_t + b_c) + (W_ch_{t-1} + b_c))c_t = f_t \odot \tilde{c}_t + i_t \odot x_t

Потом:

o_t = softsign_{scaled}((W_ox_t + b_o) + (W_oh_{t-1} + b_o))

И убрал изменение x_t (x_{t,new} = x_{t,old} теперь считай).

И ещё — теперь гейт f_t такой же как и o_t, но там (и так понятно) — своя матрица и своё смещение.

Ну и конечно же:

h_t = o_t \odot c_t

То есть я убрал \tilde{h}_t.

И только тогда моя сеть стала намного лучше (и даже быстрее!).

Плюсы и минусы моей сети

Плюсы LSWM (оригинальной):

  1. Более «большие» хвосты softsign.

  2. Легкость операций (0 экспонент).

  3. Достаточно мало параметров (нету W_c).

  4. «Self‑Attention» в LWM слою.

Минусы оригинальной LSWM:

  1. Иногда «большие» хвосты softsign'а могут вредить.

  2. x_{t,new}— это на самом деле плохое вычисление которое делает x_t слишком сильным из‑за чего сеть становится более чувствительной к рандомному сиду.

  3. Отсутствие CEC — все таки карусель постоянной ошибки важна.

У модифицированной LSWM (где есть карусель постоянной ошибки которая описана в разделе «Бенчмаркинг.» и остальные модификации) есть один минус и убирается один плюс.

Этот самый минус — это хвосты softsign.

Убирается один плюс — маленькое число параметров.

Вывод

Сделаю быстрый вывод.

Constant Error Carousel — очень важная штука, без неё никуда.

Не делай x_t слишком сильным даже если потом нормируешь его softsignом.

Softsign и его масштабированная версия (для диапазона от 0 до 1) в качестве замены tanh и sigmoid — идея рабочая.

Заменить concat суммой — тоже рабочая идея.

LWM слой — тоже рабочая идея (ведь качество не упало из‑за него, модель по‑прежнему хорошо отвечает).

P.S: Это моя первая такая статья, писал на коленке, увидите изъян в математике — пишите, грамматическую ошибку увидели — тоже пишите, потому что просто минусовать статью не даёт мне нужного фидбэка чтобы я чему то научился. Гитхаб опубликую потом...