LSWM 2.0 is all you need. ai.. ai. deep learning.. ai. deep learning. DL.. ai. deep learning. DL. gated rnn.. ai. deep learning. DL. gated rnn. gru.. ai. deep learning. DL. gated rnn. gru. lstm.. ai. deep learning. DL. gated rnn. gru. lstm. lswm.. ai. deep learning. DL. gated rnn. gru. lstm. lswm. machine learning.. ai. deep learning. DL. gated rnn. gru. lstm. lswm. machine learning. ml.

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

Прошло всего несколько дней с того момента как я выложил 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 слой делал сеть чувствительной к сиду (намного больше чем щас).

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

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

Если что датасет – мой, обучение – тоже как тогда [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)

Автор: pureooplover

Источник