Моя собственная Gated RNN: как работает? (и бенчмарки, конечно же). ai.. ai. gated rnn.. ai. gated rnn. gru.. ai. gated rnn. gru. lstm.. ai. gated rnn. gru. lstm. ml.

Всем привет!

Не давно сделал свою 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

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

68609

68880

68609

Скорость обучения (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_c(x_t + h_{t-1}) + b_c)c_t=f_t odot tilde{c}_t + i_t odot x_t

Потом:

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

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

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

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: Это моя первая такая статья, писал на коленке, увидите изъян в математике – пишите, грамматическую ошибку увидели – тоже пишите, потому что просто минусовать статью не даёт мне нужного фидбэка чтобы я чему то научился. Гитхаб опубликую потом…

Автор: pureooplover

Источник

  • Запись добавлена: 26.09.2026 в 17:52
  • Оставлено в