Снова всем привет!

Прошло всего несколько дней с того момента как я выложил LSWM [1].

В общем я протестировал и поэкспериментировал эту сеть ещё раз и нашёл СТОЛЬКО проблем, сколько даже ванильный RNN не видел.

В этой статье я попытаюсь их исправить.

Содержание.

В этой статье будет:

  1. Теория.

  2. Бенчмарки.

  3. Плюсы и минусы.

  4. Вывод.

Теория.

Начнем с теории.

Хочу сказать - тут будут те же softsign и scaled softsign [2].

И так... Начнём с первой формулы:

comb_t = x_t + h_{t-1} + n_{t-1}

Что такое n_{t-1}? Это тот же c_t, но который я прогнал через softsign. В принципе, можно записать эту формулу как:

comb_t = x_t + h_{t-1} + softsign(c_{t-1})

Но, оставим n_{t}.

Потом идут первые два гейта (обновлённые):

f_t = softsign_{scaled}(W_fcomb_t + b_f)

В принципе тот же f_t.

Теперь i_t:

i_t = softsign_{scaled}(W_icomb_t + b_i)

Как видим, я убрал два линейных слоя из прошлой статьи и заменил их одним линейным слоем.

Зачем? Качество - вверх, количество параметров - намного меньше.

Ну а вот теперь самое главное нововведение:

q_t = W_q(comb_t \odot softsign(q_{t-1})) + b_qk_t = W_k(comb_t \odot softsign(k_{t-1})) + b_kv_t = W_v(comb_t \odot softsign(v_{t-1})) + b_v

Как видим - мы прогоняем comb_t через три линейных слоя, прям как в трансформере (или в xLSTM). Но, дело не в том что мы просто "прогнали comb_t через три слоя", дело в том что из-за умножения (поэлементного) comb_t на прошлый q_t, k_t или v_t - в общем, обобщая - новый q_t становится зависим от прошлой истории q_t, k_t становится зависим от прошлой истории k_t, ну, а v_t становится зависим от прошлой истории v_t.

Это одна из причин почему LSWM 2.0 может выдерживать большие последовательности.

Что ж делать дальше?

Дальше только c_t:

c_t = f_t \odot c_{t-1} + i_t \odot (k_t \odot v_t)

Раньше [3] мы умножали i_t на обычный x_t (грех!).

Сейчас же мы берем ключ и значение нашего комбинированного значения и умножаем их (поэлементно), и вот только i_t умножаем на результат.

Сеть стала менее чувствительна к сиду и более лучше запоминать!

Потом идет долгожданный n_t:

n_t = softsign(c_t)

Нечего сверхъестественного, просто c_t прогоняем через softsign.

Потом главное (нет):

o_t = softsign_{scaled}(W_oc_t + b_o)

Я решил сделать так, чтобы o_t был зависим не от comb_t, а от самого c_t.

И это даже сделало качество лучше!

Потом идет:

h_t = o_t \odot softsign(c_t \odot q_t)

Запрос нашего comb_t умножаем на c_t, прогоняем через softsign, и конечно же умножаем o_t на этот результат.

Всё, это вся сеть.

И ещё:

Output = LayerNorm(h_t) \text{ или } H \text{ где } H = (h_0, h_1, ..., h_L)

В общем я полностью убрал LWM слой [4].

Как оказалось, LWM слой делал сеть чувствительной к сиду (намного больше чем щас).

И ещё - LSWM расшифровывается как Long-Short Working Memory.

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

И так, для начало бенчмаркинг.

Если что датасет - мой, обучение - тоже как тогда [5].

Но, для начало я замерил закон масштабирования (правильно выразился?) в 3д графике:

Результат (изображение чуть обрезано).
Результат (изображение чуть обрезано).

Вот как я тут считал:

  1. hidden dim = 128, 400 эпох.

  2. hidden dim = 2048, 1200 эпох.

  3. hidden dim = 4096, 1500 эпох.

  4. hidden dim = 128, 20000 эпох.

И ещё давненько я считал 2д график (log-log):

Результат.
Результат.

И так. Как видим, кажется, если масштабировать LSWM (если что я замерял по новой версии) - то не будет никаких особо подвохов.

Время делать бенчмарк с остальными.

Я взял тот же датасет, но проверяю модель (то есть инференс делаю) на более БОЛЬШОЙ цепочке, да.

Вот если что сам инференс:

model.eval()

with torch.no_grad():
    tokens_pool = [a, b, c]
    random_noise = []
    for _ in range(1000):
        if _ == 500:
            random_noise.extend([a])
        random_noise.extend([random.choice(tokens_pool), TOKEN_ARROW])

    chain1 = random_noise + [c, TOKEN_Q_1]
    chain2 = random_noise + [c, TOKEN_Q_2]

    test1 = torch.tensor([chain1]).to(device)
    test2 = torch.tensor([chain2]).to(device)

    pred_live = torch.argmax(model(test1), dim=1).item()
    pred_work = torch.argmax(model(test2), dim=1).item()

    print(f"need: 3 answer: {pred_live}")
    print(f"need: 12 answer: {pred_work}")

Э-э-э, ну написано довольно плохо, но оно работает.

Настроил три сети (LSWM, LSTM, GRU) под правильный размер параметров и правильные настройки биасов, и пошёл проверять.

Вот как я проверял:

  • Запускал каждую сеть 10 раз и считал сколько раз отвечала неверно и сколько раз верно.

  • Та, которая показала себя лучше всего (то есть верных ответов больше чем неверных чем у остальных) - та и победила.

Вот результаты:

Сеть

Верных

Неверных

Параметров

LSWM (2.0)

5

5

101632

GRU

3

7

101632

LSTM

4

6

101672

Лосс и качество убрал - потому что я лентяй и забыл считать хотя бы среднее между всеми 10 запусками, но в принципе у всех там 100%-93% качество в основном.

И ещё - я не могу гарантировать что эта статистика верных и неверных всегда будет совпадать с таблицей, но примерно так всегда будет.

Ну, а теперь по самой таблице:

  1. LSWM 2.0 переиграла всех.

  2. GRU, соответственно, хуже всех.

  3. LSTM "на втором месте".

Это означает что LSWM 2.0 может конкурировать.

Плюсы и минусы.

Плюсы:

  1. (В теории) Больше запоминает.

  2. Чувствительность к сиду намного меньше (но ещё есть).

  3. Обучается не так уж и не медленно.

  4. o_t "принимает решение" на основе существующей памяти (c_t), что может очень хорошо сказаться.

Минусы:

  1. Всё чувствительность к сиду присутствует (этот минус можно и не писать, я его считай описал в плюсах...).

  2. softsign иногда своими "хвостами" только мешает (но я пока такого не видел).

  3. Сеть к сожалению всё равно с трудом проходит задачи по типу Needle in the haystack.

В общем, я считаю что эта версия LSWM достаточно конкурентноспособная, но всё же пока что тестирую ещё.

Вывод.

Засунуть Q, K, V в последовательную Gated RNN - штука рабочая.

Суммирование вместо конкатенирования - очень рабочая идея.

softsign - очень хорошо.

Но стоит и отметить, что нельзя всё пихать в одну кучу, что я и подтвердил на примере LWM.

P.S: так как проблемы с Github'ом до сих пор - держите код LSWM:

Код.
import torch
import torch.nn as nn
import torch.optim as optim
import torch.nn.functional as F
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.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_i = nn.Linear(d_model, d_model).to(device)
        self.W_o = nn.Linear(d_model, d_model).to(device)

        with torch.no_grad():
            self.W_f.bias.fill_(3.0)
            self.W_i.bias.fill_(0.0)
            self.W_o.bias.fill_(0.0)

        self.W_q = nn.Linear(d_model, d_model).to(device)
        self.W_k = nn.Linear(d_model, d_model).to(device)
        self.W_v = nn.Linear(d_model, d_model).to(device)

        self.norm = nn.LayerNorm(d_model)

    def softsign_scaled(self, x):
        return (F.softsign(x) + 1.0) / 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)
        c_t = torch.zeros(batch_size, self.d).to(device)

        q_t = torch.ones(batch_size, self.d).to(device)
        k_t = torch.ones(batch_size, self.d).to(device)
        v_t = torch.ones(batch_size, self.d).to(device)

        n_t = torch.zeros(batch_size, self.d).to(device)

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

            combined = h_t + x_t + n_t

            f_t = self.softsign_scaled(self.W_f(combined))
            i_t = self.softsign_scaled(self.W_i(combined))

            q_t = self.W_q(combined * F.softsign(q_t))
            k_t = self.W_k(combined * F.softsign(k_t))
            v_t = self.W_v(combined * F.softsign(v_t))

            c_t = f_t * c_t + i_t * (k_t * v_t)
            n_t = F.softsign(c_t)

            o_t = self.softsign_scaled(self.W_o(c_t))

            h_t = o_t * F.softsign(c_t * q_t)

        return self.norm(h_t)

Комментарии (5)


  1. pureooplover Автор
    01.10.2026 11:41

    Если нашли какую-то ошибку, недочёт или что либо ещё - сообщайте пожалуйста.


    1. ToxaBes
      01.10.2026 11:41

      При прочтении вашей работы (этой и предыдущих частей) меня не покидает стойкое чувство дежавю. Как будто я уже читал как это кто-то делал или разбирал когда-то, но я не могу ни вспомнить, ни найти подтверждение этому.

      Я ни в коем случае не имею ввиду, что ваша работа вторична или что-то такое, я о том что походу у меня когнитивное искажение в конкретно этой теме. Вы, пожалуйста, продолжайте, интересно что это даст дальше. Мне почему-то кажется, что residual connections даст больше эффект чем mLSTM.


      1. pureooplover Автор
        01.10.2026 11:41

        Мне почему-то кажется, что residual connections даст больше эффект чем mLSTM

        Я имел ввиду объединить их...

        Ну типа:

        1. LSWM

        2. Residual connection (то есть то что пошло на вход мы прибавляем к результату слоя)

        3. mLSTM.

        4. Выход.


        1. ToxaBes
          01.10.2026 11:41

          Я в том смысле, что использование mLSTM скорее всего все ухудшит, а сам Residual connection для с вязи с чем-то другим - нет. Но могу ошибаться.


          1. pureooplover Автор
            01.10.2026 11:41

            А-а-а, теперь понял!