Мой блог
Архитектура LSWM 2.0: дорабатываем нейросеть и побеждаем классические рекуррентные модели
С момента публикации первой версии архитектуры LSWM прошло совсем немного времени, но повторные тесты и эксперименты вскрыли множество скрытых архитектурных проблем. В этой статье я подробно разберу, как мне удалось исправить ключевые недостатки сети, переработать математические формулы и добиться превосходства над стандартными LSTM и GRU в бенчмарках.
Проведенные тесты показывают, что LSWM 2.0 уверенно выдерживает длинные последовательности, снижает чувствительность к сидам и демонстрирует отличные результаты на практике. Если вам интересна практическая разработка нейросетей и оптимизация рекуррентных архитектур, этот материал даст исчерпывающее техническое руководство.
Теоретические основы и математика LSWM 2.0
В основе обновленной архитектуры лежат функции активации softsign и её масштабированная модификация scaled softsign. При проектировании новой версии я полностью пересмотрел структуру гейтов, чтобы повысить общую стабильность и качество обработки информации.
Обновленные гейты и линейные слои
Первое важное изменение коснулось структуры линейных слоев. В прошлой версии использовалось два линейных слоя, которые я успешно объединил в один. Это позволило одновременно поднять качество работы модели и существенно сократить общее количество обучаемых параметров.
Самое главное нововведение заключается в интеграции принципов, заимствованных из трансформеров и архитектуры xLSTM. Мы последовательно пропускаем данные через три линейных слоя, вычисляя запросы, ключи и значения (Q, K, V). За счет поэлементного умножения на предыдущие состояния новый шаг становится глубоко зависимым от всей предшествующей истории.
Отказ от LWM-слоя и стабилизация памяти
Предыдущие версии страдали от избыточной чувствительности к сидам из-за использования устаревшего LWM-слоя. В версии 2.0 я полностью избавился от этого элемента, что моментально сделало сеть менее зависимой от случайных инициализаций и улучшило способность запоминать контекст.
Сама аббревиатура LSWM расшифровывается как Long-Short Working Memory. Интеграция механизмов Q, K, V в последовательную Gated RNN и замена конкатенации на суммирование открывают отличные возможности для эффективной работы с длинными контекстами без усложнения логики.
Бенчмаркинг и сравнительное тестирование моделей
Для оценки возможностей новой архитектуры я подготовил собственный датасет и провел серию комплексных тестов. Сначала я замерил закон масштабирования в трехмерном пространстве, варьируя размер скрытого слоя (hidden dim от 128 до 4096) и количество эпох обучения.
Анализ лог-лог графиков показал, что масштабирование LSWM 2.0 происходит стабильно и без неприятных сюрпризов. На следующем этапе я перешел к прямому сравнению с другими популярными рекуррентными сетями.
Сравнение LSWM 2.0, LSTM и GRU
Я настроил три различные модели — LSWM, LSTM и GRU — под идентичный объем параметров и корректные настройки смещений (biases). Каждая сеть запускалась по десять раз на увеличенной цепочке токенов для оценки инференса.
В ходе тестов фиксировалось количество правильных и ошибочных ответов. Несмотря на то, что качество у всех моделей держалось на высоком уровне в диапазоне от 93% до 100%, по количеству успешных ответов LSWM 2.0 уверенно заняла первое место, оставив LSTM и GRU позади.
Плюсы, минусы и практические выводы
Любая экспериментальная архитектура имеет свои сильные и слабые стороны. Среди ключевых преимуществ LSWM 2.0 стоит выделить сниженную чувствительность к сидам, высокую скорость обучения и способность принимать решения на основе накопленной рабочей памяти.
К недостаткам можно отнести особенности работы функций активации softsign, чьи хвосты теоретически могут мешать обучению, хотя на практике я с этим не сталкивался. Кроме того, задачи типа Needle in the haystack даются сети все еще с трудом.
Тем не менее, текущая версия демонстрирует высокую конкурентоспособность. Эксперименты продолжаются, а для тех, кто хочет протестировать разработку самостоятельно, ниже приведена готовая реализация сети на PyTorch.
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)
