Перейти к содержанию
Публикация AiManual

Своя Gated RNN вместо LSTM и GRU: разбор математики LSWM, тестов и бенчмарков

Автор LSWM заменил tanh и сигмоиду на softsign, убрал CEC и заменил конкатенацию сложением. По его замерам на CPU модель дала 100% точности на multi-hop branchi

Коротко

Что будет в материале

  1. 01

    Что такое LSWM и зачем понадобилась своя gated RNN

  2. 02

    Математика LSWM: softsign вместо tanh, scaled softsign вместо сигмоиды и сложение вместо конкатенации

  3. 03

    Два этапа памяти: SWM и LWM с механизмом, близким к self-attention

  4. 04

    Бенчмарки: LSWM против LSTM и GRU на задаче multi-hop branching

Что такое LSWM и зачем понадобилась своя gated RNN

LSWM (Long-Short Working Memory) - gated рекуррентная сеть, которую собрал автор и разобрал в статье на Habr. От LSTM, GRU и похожих архитектур её отличает математика: вместо tanh и сигмоиды работают softsign и его масштабированная версия, а конкатенация входов заменена сложением.

Внутри два этапа. Первый, SWM, это краткосрочная рабочая память с forget-, input- и output-гейтами, как в LSTM. Второй, LWM, долгосрочная рабочая память: результат умножается на матрицу, которую автор осторожно называет «матрицей внимания», и сам же сомневается, уместен ли термин. Механизм близок к self-attention, потому что каждый элемент последовательности получает вес с учётом остальных.

Практический результат такой: на задаче multi-hop branching LSWM дала 100% точности против 96.9% у GRU и 71.9% у LSTM при 400 эпохах обучения. Цена - время: около 8 секунд против 4 секунд у базовых моделей. Замеры сделаны на CPU в Google Colab, и автор сразу оговаривает, что не умеет замерять так, как это принято в научных работах, и замеряет «как может».

Разбор ценен не цифрами, а тем, что создатель не скрывает проблем: чувствительность к seed, отсутствие Constant Error Carousel, слишком сильное влияние softsign и последующие правки формул. Дальше по порядку: математика, устройство двух этапов памяти, замеры и ограничения.

Математика LSWM: softsign вместо tanh, scaled softsign вместо сигмоиды и сложение вместо конкатенации

В классической LSTM-ячейке работают две нелинейности: tanh с диапазоном от -1 до 1 и сигмоида от 0 до 1. Гейты решают, сколько информации пропустить, tanh формирует кандидата на новое состояние памяти. В LSWM обе функции заменены на семейство softsign, и есть третье отличие: там, где обычно вход и скрытое состояние склеиваются в один вектор, здесь они складываются.

Почему softsign, а не tanh: диапазоны и поведение

Обычный softsign считается как x / (1 + |x|). Диапазон значений тот же, что у tanh: от -1 до 1. Различается форма кривой. Softsign подходит к границам медленнее, его производная убывает как 1 / (1 + |x|)², тогда как у tanh она падает почти до нуля уже при |x| около 3. Для рекуррентной сети это значит, что ненулевой градиент сохраняется на большем диапазоне входов.

Автор не заявляет, что softsign объективно лучше tanh. По его данным видно другое: функция оказалась капризной. Оценить отдельный вклад замены активации по опубликованным цифрам нельзя, потому что в тот же замер попали двухэтапная память, сложение вместо конкатенации и отказ от CEC.

Scaled softsign: как получить диапазон от 0 до 1 без сигмоиды

Гейтам нужен диапазон от 0 до 1, чтобы работать множителями: ноль стирает информацию, единица пропускает её целиком. Обычный softsign для этого не подходит, он уходит в минус. Автор сделал масштабированную версию и объясняет задачу так:

«Эта формула мне нужна для замены сигмоиды (сигмоида выдает диапазон от 0 до 1, а обычный софтсайн - от -1 до 1, а мне нужно было от 0 до 1).»

В формуле участвует обучаемый параметр, который автор изначально ставит на 2: «умножаем на обучаемый параметр (я его на 2 с начало ставлю, вроде так лучше по качеству и обобщению)». Масштаб меняет крутизну кривой, и сеть подбирает его при обучении вместо жёстко заданной константы.

Порядок применения функций критичен. Автор пробовал переставить их: обычный softsign в вычислении кандидата, scaled softsign в output-гейте. Точность просела до 23%. Размещение нелинейностей подбиралось экспериментально, а не выводилось из теории.

Сложение вместо конкатенации: что это меняет

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

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

Два этапа памяти: SWM и LWM с механизмом, близким к self-attention

Ячейка LSWM считает состояние за два прохода: сначала короткая память, потом длинная.

SWM: forget, input и output гейты в краткосрочной памяти

Первый этап, SWM, повторяет гейтовый набор LSTM. Forget-гейт решает, что стереть из предыдущего состояния, input-гейт определяет, что записать, output-гейт - что отдать на выход. Отличия в нелинейностях и в сложении вместо конкатенации. Автор описывает этот этап как первый в вычислении и отдельно предупреждает, что долгосрочной памяти здесь ещё нет: «Это первый этап вычисления в моей Gated RNN. Возможно вы спросите - а где же долгосрочная (Long-Term) память? Ну, вот щас покажу.»

LWM: как работает долгосрочная память и почему это похоже на self-attention

Второй этап умножает результат SWM на матрицу, где каждый элемент получает вес с учётом остальных, и делит результат, чтобы градиенты не взорвались при обратном проходе. Как это описывает сам автор: «Умножаем на эту самую "матрицу внимания" чтобы сделать размерность правильной и делим на [...] чтобы не взорвать градиенты. Всё, это вся долгосрочная память.»

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

Бенчмарки: LSWM против LSTM и GRU на задаче multi-hop branching

Условия эксперимента: CPU, Google Colab, 400 эпох

Замеры сделаны на CPU в Google Colab, обучение шло 400 эпох на задаче multi-hop branching. Такие задачи проверяют, умеет ли сеть удерживать информацию и возвращаться к ней через несколько шагов, поэтому они чувствительны к качеству памяти модели.

Статус цифр автор оговаривает заранее: он не умеет замерять так, как это делается в научных работах, и замеряет «как может». Разброс по seed в разборе не приводится, а архитектура к seed чувствительна. Таблицу ниже стоит читать как ориентир, а не как воспроизводимый бенчмарк.

Результаты: точность и время обучения

АрхитектураТочность, 400 эпохВремя обучения
LSWM100%около 8 секунд
GRU96.9%около 4 секунд
LSTM71.9%около 4 секунд

Отрыв от GRU - 3.1 процентного пункта при вдвое большем времени обучения. Разрыв с LSTM больше, но обе цифры получены в одном прогоне без усреднения по seed, поэтому разница между 71.9% и 100% отражает в том числе разброс обучения. Чтобы утверждать, что дело именно в архитектуре, нужны серии запусков. Цифры и условия взяты из замеров автора.

Проблемы и ограничения LSWM: чувствительность к seed, отсутствие CEC и влияние softsign

Автор перечисляет слабые места сам, и этот список важнее таблицы с точностью.

Почему отсутствие Constant Error Carousel меняет поведение сети

CEC (Constant Error Carousel) в LSTM сохраняет градиент вдоль потока состояния памяти: ошибка идёт по этой линии без умножения на производные нелинейностей, поэтому сеть способна запоминать информацию на сотни шагов. Автор убрал механизм сознательно, ещё на старте: «где CEC? Ответ прост - я решил убрать эту всю мишуру (ладно, это не мишура) ещё на старте». Оговорка в скобках честно признаёт, что вещь это не лишняя.

Без CEC градиент по времени проходит через обычные произведения матриц и производных активаций, а такое произведение склонно либо затухать, либо расти. Насколько сильно это ограничивает LSWM, по опубликованным замерам не увидеть: нужны отдельные тесты на длинных зависимостях.

Чувствительность к seed и нестабильность результатов

При другом начальном заполнении весов результат может отличаться, и автор называет это отдельной проблемой архитектуры. Конкретных цифр разброса в разборе нет, поэтому 100% точности корректно читать как один из возможных результатов, а не как гарантированное свойство модели. Для практики это означает усреднение по нескольким запускам.

Слишком сильное влияние softsign и последующие правки формул

Softsign насыщается медленнее tanh и пропускает дальше более крупные значения. В LSWM это влияние оказалось чрезмерным, и формулы пришлось менять уже после первых экспериментов. Масштаб чувствительности виден по обратному примеру: перестановка обычного и масштабированного softsign по местам обрушила точность до 23%. Любая правка нелинейностей здесь требует переобучения и повторных замеров. О проблемах создатель пишет прямо в исходном разборе.

Что из LSTM действительно критично, а что можно заменить

История LSWM позволяет разложить LSTM на элементы и проверить, что переживает замену, а что нет.

  • Гейтовый каркас. Forget-, input- и output-гейты сохранены в SWM, и сеть работает. Идея управляемого пропуска информации через множители с диапазоном от 0 до 1 осталась нетронутой, изменилась только функция, которая эти множители считает.
  • Нелинейности. Замена tanh на softsign и сигмоиды на scaled softsign прошла, но оказалась самой хрупкой частью. Неверное размещение функций дало падение до 23%. Взаимозависимость нелинейностей внутри ячейки выше, чем кажется по формулам.
  • CEC. Единственный элемент, который автор убрал полностью, и главный теоретический риск. Именно CEC отвечает за сохранение градиента во времени.
  • Конкатенация. Замена сложением не сломала обучение, но лишила ячейку раздельного учёта входа и скрытого состояния.

Картина по элементам такая: гейтовая схема и множители с диапазоном от 0 до 1 устойчивы к заменам, а конкретные функции активации и способ объединения входов подбираются под задачу и требуют проверки. Убирать CEC без другого механизма сохранения градиента - самый рискованный шаг из четырёх.

Практические выводы: стоит ли использовать LSWM и как экспериментировать с RNN самостоятельно

LSWM остаётся исследовательским проектом одного автора, а не готовой заменой LSTM и GRU. Она выиграла на одной синтетической задаче, обучалась вдвое дольше базовых моделей и требует аккуратной настройки нелинейностей. Для рабочего проекта разумнее брать проверенные варианты: GRU или LSTM, если нужна рекуррентность, Transformer, если важны длинные зависимости и параллельное обучение. LSWM интересна как материал для изучения того, как устроены gated RNN изнутри.

Если хочется повторить такой эксперимент, порядок действий примерно такой:

  1. Возьмите задачу с автоматически проверяемым ответом, как multi-hop branching. Тогда точность считается без ручной разметки.
  2. Зафиксируйте seed и прогоняйте каждую конфигурацию несколько раз. Без этого легко перепутать удачный запуск с работающей идеей.
  3. Меняйте по одному элементу: сначала одну нелинейность, потом способ объединения входов. Иначе непонятно, что дало прирост, а что его съело.
  4. Замеряйте время обучения вместе с точностью. Прирост в 3 процентных пункта при удвоенном времени подходит не для каждой задачи.
  5. Держите под рукой базовую модель. GRU дала 96.9% за вдвое меньшее время, и обойти её сложнее, чем кажется.

Логика изоляции переменных работает и за пределами рекуррентных сетей. В проекте tiny-sparse-lab автор отдельно проверял, можно ли вынести знания из весов модели в разреженную память, и прогнал серию контролей, прежде чем делать выводы: разбор этих экспериментов полезен как пример инфраструктуры проверки для гипотезы, которая пока не подтверждена.

Самый быстрый способ проверить всё это на себе: взять готовую реализацию LSTM из своей библиотеки, написать рядом свою ячейку с softsign вместо tanh и прогнать оба варианта на одной задаче с одинаковым seed и одинаковым числом эпох. Разница в точности и времени скажет больше, чем любое описание формул.

Подписаться на канал