Вход на сайт

Просмотр новости

Найдите то, что Вас интересует

Моя собственная Gated RNN: как работает? (и бенчмарки, конечно же)

Дата публикации: 26-09-2026 17:52:14

Всем привет!Не давно сделал свою Gated RNN, то есть с другой математикой, ни как у LSTM, GRU и подобного.Я хочу (для вас) разобрать её теоретически, практически, замерить (я не умею замерять так что буду замерять как могу), ну и конечно же расскажу плюсы и минусы. Читать далее

Основное содержимое страницы с новостью.

Время на прочтение8 мин

Охват и читатели565

Всем привет!

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

Схожие новости

#Наименование новостиТональностьИнформативностьДата публикации
1Параллельность RNN?07.6607-06-2026
2Ваш бэктест нейронки врёт на ±10 процентных пунктов08.9912-08-2026
3Ежедневный Хабр: 9 интересных публикаций каждый день012.2106-08-2026
4Как я смог запустить 32b модель и сделать ее умнее в два раза на 305005.9529-07-2026
5[Перевод] Как на самом деле работают LLM0707-07-2026
6Математика stop‑loss в Telegram Ads. Когда выключать безрезультатное объявление09.9310-08-2026
7Нейросеть-автопилот вместо 400 Playwright-тестов0707-07-2026
8Нейро сети для самых маленьких. Часть первая (которая после нулевой). Удобство в прокрустовом ложе оптимизации07.5401-07-2026
9Как желание быстрее читать чужой код превратилось в войну с недетерминизмом LLM0528-06-2026
10Рациональная параметризация как метод устранения численных погрешностей при построении графиков-16.4426-09-2026

Классификация: Наука. Схожих патентов: 0. Схожих новостей: 10. Тональность: 0. Информативность: 5.62. Источник: habr.com.