Загрузка 0
ПОДЕЛИТЬСЯ

Мой блог

Листай вниз

Архитектура LSWM 2.0: дорабатываем нейросеть и побеждаем классические рекуррентные модели

Архитектура LSWM 2.0: дорабатываем нейросеть и побеждаем классические рекуррентные модели

С момента публикации первой версии архитектуры LSWM прошло совсем немного времени, но повторные тесты и эксперименты вскрыли множество скрытых архитектурных проблем. В этой статье я подробно разберу, как мне удалось исправить ключевые недостатки сети, переработать математические формулы и добиться превосходства над стандартными LSTM и GRU в бенчмарках.

Проведенные тесты показывают, что LSWM 2.0 уверенно выдерживает длинные последовательности, снижает чувствительность к сидам и демонстрирует отличные результаты на практике. Если вам интересна практическая разработка нейросетей и оптимизация рекуррентных архитектур, этот материал даст исчерпывающее техническое руководство.

Формула масштабированного softsign для нейросети
Применение функции scaled softsign для стабилизации вычислений внутри гейтов.
Двумерный log-log график масштабирования модели
Двумерный логарифмический график для оценки масштабируемости обновленной сети.
Код реализации LSWM на PyTorch
Фрагмент программного кода архитектуры LSWM с использованием библиотек PyTorch.
Сравнительная таблица результатов тестов LSWM, GRU и LSTM
Таблица сравнения точности и количества параметров LSWM 2.0, GRU и LSTM.

Теоретические основы и математика LSWM 2.0

В основе обновленной архитектуры лежат функции активации softsign и её масштабированная модификация scaled softsign. При проектировании новой версии я полностью пересмотрел структуру гейтов, чтобы повысить общую стабильность и качество обработки информации.

Реклама
Интерфейс опроса разработчиков о перспективах архитектуры
Голосование пользователей по поводу интеграции дополнительных модулей в LSWM.
График функции активации в нейросети
Визуализация поведения функции активации softsign на внутренних слоях сети.
Схема линейных слоев и гейтов LSWM 2.0
Схема взаимодействия линейных слоев и гейтов в новой версии модели.
Пример кода инференса модели для проверки последовательностей
Код проверки модели на больших цепочках данных во время инференса.

Обновленные гейты и линейные слои

Первое важное изменение коснулось структуры линейных слоев. В прошлой версии использовалось два линейных слоя, которые я успешно объединил в один. Это позволило одновременно поднять качество работы модели и существенно сократить общее количество обучаемых параметров.

Реклама
Структура весов и биасов в инициализации сети
Настройка начальных параметров смещения для стабилизации обучения.
Визуализация скрытых состояний рекуррентной сети
График распределения скрытых состояний и векторов памяти во времени.
Сравнение работы конкатенации и суммирования в гейтах
Иллюстрация преимущества замены конкатенации на поэлементное суммирование.

Самое главное нововведение заключается в интеграции принципов, заимствованных из трансформеров и архитектуры xLSTM. Мы последовательно пропускаем данные через три линейных слоя, вычисляя запросы, ключи и значения (Q, K, V). За счет поэлементного умножения на предыдущие состояния новый шаг становится глубоко зависимым от всей предшествующей истории.

Тестирование устойчивости модели к сидам
Графики разброса результатов при различных инициализациях случайных чисел.
График распределения лосса при обучении LSWM 2.0
График падения функции потерь (loss) в процессе обучения модели.
Схема интеграции механизмов внимания QKV в RNN
Схема работы запросов, ключей и значений внутри последовательной модели.
Тестовый стенд для проверки длинных цепочек токенов
Проверка способности сети удерживать контекст на расширенных цепочках.

Отказ от LWM-слоя и стабилизация памяти

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

Пример работы токенизатора в экспериментальном пайплайне
Подготовка и подача токенов на входembedding-слоя нейросети.
Настройка нормализации слоев LayerNorm в LSWM
Нормализация выходных данных для предотвращения взрыва градиентов.
Анализ ошибок моделей в сравнительном тесте
Сравнение частоты ошибок LSWM 2.0 с классическими архитектурами.

Сама аббревиатура LSWM расшифровывается как Long-Short Working Memory. Интеграция механизмов Q, K, V в последовательную Gated RNN и замена конкатенации на суммирование открывают отличные возможности для эффективной работы с длинными контекстами без усложнения логики.

Общий вид архитектурной схемы LSWM 2.0
Полная архитектурная схема обновленной рекуррентной нейросети.

Бенчмаркинг и сравнительное тестирование моделей

Для оценки возможностей новой архитектуры я подготовил собственный датасет и провел серию комплексных тестов. Сначала я замерил закон масштабирования в трехмерном пространстве, варьируя размер скрытого слоя (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)
01.