Мой блог
Разработка собственной Gated RNN: теория, архитектура, практика и бенчмарки
Приветствую читателей блога Sergey Bagrov! Недавно я занялся разработкой своей Gated RNN с принципиально иной математикой, непохожей на классические архитектуры вроде LSTM или GRU. В этой публикации я подробно разберу теоретическую базу своей модели, поделюсь практической реализацией, проведу доступные бенчмарки и честно рассмотрю все сильные и слабые стороны получившегося решения.
Это моя первая попытка написать материал, близкий к научному формату, поэтому здесь могут встречаться определенные огрехи. Ниже приведена структура статьи: сначала теория и формулы, затем практические тесты на сложных задачах, бенчмарки на CPU, анализ плюсов с минусами и итоговые выводы.
Теоретическая база и архитектура сети
Начнем с математического фундамента и нестандартных решений, которые я заложил в основу. Для замены стандартных функций активации мне потребовались модифицированные варианты. Обычный сигмоид выдает диапазон от 0 до 1, а стандартный софтсайн — от -1 до 1, но для моих задач требовался строгий диапазон от 0 до 1. Для этого я применил формулу масштабирования softsign, где результат делится на его модуль плюс единица.
Первый этап вычислений в моей сети получил название SWM (Short Working Memory). Для реализации своеобразного насыщения, улучшающего обобщение модели, я использую специальное преобразование: старое состояние суммируется с входными данными, пропускается через масштабированный софтсайн и умножается на обучаемый параметр (на старте я задаю его равным двум для лучшего качества).
Гейты и вычисление кандидата
Гейт забывания (forget gate) устроен схоже с механизмом в обычном LSTM, однако вместо конкатенации я используем операцию сложения и собственную функцию активации. Далее задействуется гейт ввода (input gate): входные данные и скрытое состояние проходят через два линейных слоя со смещением, результаты суммируются и пропускаются через масштабированный софтсайн.
Вычисление кандидата строится на поэлементном перемножении гейтов и входов с последующей обработкой через scaled softsign. Хотя я сам имею определенные сомнения насчет такого подхода, на практике для моих задач это оказался наиболее работоспособный вариант. Попытки поменять местами функции активации в кандидате и выходном гейте приводили к резкому падению качества модели.
Долгосрочная память и итоговая цепочка
Многие читатели могут задаться вопросом о наличии долгосрочной памяти (Long-Term Memory) и карусели постоянной ошибки (CEC). На старте экспериментов я осознанно убрал эту «мишуру», и сеть заработала, поэтому на начальном этапе все осталось именно так. Второй этап архитектуры — LWM (Long Working Memory) — работает с размерностью скрытого состояния, умножая матрицу состояний для получения аналога механизма внимания.
Полученная матрица делится на квадратный корень из размерности модели, чтобы избежать взрыва градиентов. Если данный слой является финальным в сети, далее считается сумма каждой строки полученной матрицы в единый список. Завершает структуру SWM и LWM стандартный слой LayerNorm, после чего может следовать классификатор. Всю эту архитектуру я объединил под общим названием LSWM (Long-Short Working Memory).
Практическая проверка и сложные тесты
Для проверки возможностей сети я выбрал достаточно сложную задачу, известную как Multi-hop branching. Обычный multi-hop проверяет цепочки вроде равенства переменных, но для моей модели это оказалось слишком простой задачей. Ветвящийся вариант multi-hop проверяет корректность задания переменных в начале и конце последовательности с помощью специальных инференсов.
Во время экспериментов на втором тесте (экстраполяция длины последовательности) моя сеть регулярно допускала ошибки. Первоначально я грешил на слой LWM и перепробовал множество вариантов, включая нормализацию и матрицы проекций Q, K, V. Однако решение оказалось проще — я забыл добавить output gate. Час экспериментов с выходным гейтом позволил подобрать нужную формулу, после чего сеть успешно справилась со всеми ветвящимися цепочками.
Дополнительные тесты с добавлением «мусорных» токенов показали неожиданно высокую стабильность модели. Итоговые проверки подтвердили, что LSWM демонстрирует отличные результаты наряду с другими архитектурами, хотя впереди еще предстоит тестирование на задачах генерации текста.
Бенчмаркинг производительности
Поскольку лимиты на графический ускоритель в Google Colab были исчерпаны, все бенчмарки я проводил на центральном процессоре (CPU), используя первый ветвящийся multi-hop тест. Сравнение проводилось с классическими моделями LSTM и GRU при одинаковых гиперпараметрах.
По итогам 400 эпох обучения моя LSWM обошла конкурентов по точности, достигнув абсолютного результата, в то время как LSTM и GRU показали высокую скорость обучения. Однако все три модели выдавали одинаково правильные итоговые ответы на тестовых цепочках. Дальнейший анализ выявил чувствительность модели к случайному сиду (seed), поэтому в дальнейшем я оптимизировал математику LSWM, убрав лишнее умножение на обучаемый параметр и повысив общую стабильность.
Плюсы и минусы собственной архитектуры
Оригинальная версия LSWM обладает рядом весомых преимуществ, но не лишена недостатков, которые проявились в ходе экспериментов и бенчмарков.
Достоинства и недостатки
К плюсам оригинальной сети можно отнести легкость вычислительных операций без экспонент, небольшое количество параметров за счет отсутствия матрицы $W_c$, а также встроенные свойства внимания в LWM-слое. К минусам относятся потенциальный вред от больших хвостов функции softsign и повышенная чувствительность к случайному сиду из-за особенностей масштабирования.
В модифицированной версии LSWM, где была возвращена карусель постоянной ошибки (CEC), часть плюсов и минусов перераспределилась. Модель стала надежнее, но при этом увеличилось число параметров и вернулся эффект влияния хвостов активации.
Заключение
Подводя итог проделанной работе, можно сделать несколько важных практических выводов. Карусель постоянной ошибки (Constant Error Carousel) остается критически важным элементом для стабильной работы рекуррентных сетей. Использование softsign и его масштабированных модификаций в качестве замены tanh и sigmoid доказало свою полную жизнеспособность, равно как и замена конкатенации на простое сложение. LWM-слой также отлично вписался в архитектуру, не ухудшив качество ответов. Если вы заметите неточности в математике или грамматике, обязательно пишите в комментариях — конструктивная критика помогает развиваться дальше!
