Снова всем привет!
Прошло всего несколько дней с того момента как я выложил LSWM [1].
В общем я протестировал и поэкспериментировал эту сеть ещё раз и нашёл СТОЛЬКО проблем, сколько даже ванильный RNN не видел.
В этой статье я попытаюсь их исправить.
Содержание.
В этой статье будет:
Теория.
Бенчмарки.
Плюсы и минусы.
Вывод.
Теория.
Начнем с теории.
Хочу сказать - тут будут те же softsign и scaled softsign [2].
И так... Начнём с первой формулы:
Что такое ? Это тот же
, но который я прогнал через softsign. В принципе, можно записать эту формулу как:
Но, оставим .
Потом идут первые два гейта (обновлённые):
В принципе тот же .
Теперь :
Как видим, я убрал два линейных слоя из прошлой статьи и заменил их одним линейным слоем.
Зачем? Качество - вверх, количество параметров - намного меньше.
Ну а вот теперь самое главное нововведение:
Как видим - мы прогоняем через три линейных слоя, прям как в трансформере (или в xLSTM). Но, дело не в том что мы просто "прогнали
через три слоя", дело в том что из-за умножения (поэлементного)
на прошлый
,
или
- в общем, обобщая - новый
становится зависим от прошлой истории
,
становится зависим от прошлой истории
, ну, а v_t становится зависим от прошлой истории
.
Это одна из причин почему LSWM 2.0 может выдерживать большие последовательности.
Что ж делать дальше?
Дальше только :
Раньше [3] мы умножали на обычный
(грех!).
Сейчас же мы берем ключ и значение нашего комбинированного значения и умножаем их (поэлементно), и вот только умножаем на результат.
Сеть стала менее чувствительна к сиду и более лучше запоминать!
Потом идет долгожданный :
Нечего сверхъестественного, просто прогоняем через softsign.
Потом главное (нет):
Я решил сделать так, чтобы был зависим не от
, а от самого
.
И это даже сделало качество лучше!
Потом идет:
Запрос нашего умножаем на
, прогоняем через softsign, и конечно же умножаем o_t на этот результат.
Всё, это вся сеть.
И ещё:
В общем я полностью убрал LWM слой [4].
Как оказалось, LWM слой делал сеть чувствительной к сиду (намного больше чем щас).
И ещё - LSWM расшифровывается как Long-Short Working Memory.
Бенчмаркинг.
И так, для начало бенчмаркинг.
Если что датасет - мой, обучение - тоже как тогда [5].
Но, для начало я замерил закон масштабирования (правильно выразился?) в 3д графике:

Вот как я тут считал:
hidden dim = 128, 400 эпох.
hidden dim = 2048, 1200 эпох.
hidden dim = 4096, 1500 эпох.
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% качество в основном.
И ещё - я не могу гарантировать что эта статистика верных и неверных всегда будет совпадать с таблицей, но примерно так всегда будет.
Ну, а теперь по самой таблице:
LSWM 2.0 переиграла всех.
GRU, соответственно, хуже всех.
LSTM "на втором месте".
Это означает что LSWM 2.0 может конкурировать.
Плюсы и минусы.
Плюсы:
(В теории) Больше запоминает.
Чувствительность к сиду намного меньше (но ещё есть).
Обучается не так уж и не медленно.
"принимает решение" на основе существующей памяти (
), что может очень хорошо сказаться.
Минусы:
Всё чувствительность к сиду присутствует (этот минус можно и не писать, я его считай описал в плюсах...).
softsign иногда своими "хвостами" только мешает (но я пока такого не видел).
Сеть к сожалению всё равно с трудом проходит задачи по типу Needle in the haystack.
В общем, я считаю что эта версия LSWM достаточно конкурентноспособная, но всё же пока что тестирую ещё.
Вывод.
Засунуть Q, K, V в последовательную Gated RNN - штука рабочая.
Суммирование вместо конкатенирования - очень рабочая идея.
- очень хорошо.
Но стоит и отметить, что нельзя всё пихать в одну кучу, что я и подтвердил на примере 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)

