Vanilla RNN в ConvNetJS

Vanilla RNN (Recurrent Neural Network) представляет собой базовую рекуррентную нейронную сеть, способную обрабатывать последовательные данные. В ConvNetJS реализация RNN ориентирована на простоту и наглядность, позволяя экспериментировать с временными рядами, текстом и другими последовательностями. Основные компоненты Vanilla RNN включают входной слой, скрытый слой с рекуррентными связями и выходной слой.

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

Основные классы и структуры данных

ConvNetJS предоставляет несколько ключевых классов для работы с RNN:

  • RNN — основной объект рекуррентной сети. Он управляет структурой сети, параметрами весов и обучением.
  • Layer — слой сети, может быть входным, скрытым или выходным. В контексте RNN скрытый слой содержит рекуррентные веса.
  • Vol — структура для хранения данных и градиентов. Каждый элемент последовательности представлен как объект Vol, обеспечивающий удобный доступ к данным для прямого и обратного прохода.

Весовые матрицы скрытого слоя включают:

  • Wx — веса для входного слоя.
  • Wh — рекуррентные веса для предыдущего скрытого состояния.
  • b — вектор смещений.

Прямой проход (Forward Pass)

Прямой проход в Vanilla RNN вычисляется пошагово по времени:

  1. Инициализация скрытого состояния h_0 обычно нулями.

  2. Для каждого элемента последовательности x_t выполняется:

    [ h_t = (W_x x_t + W_h h_{t-1} + b)]

  3. Выход сети y_t формируется через линейное преобразование скрытого состояния:

    [ y_t = W_y h_t + b_y]

Функция активации tanh обеспечивает нелинейность и ограничивает значения скрытого состояния в диапазоне ([-1,1]), предотвращая взрыв значений на первых шагах обучения.

В ConvNetJS прямой проход реализован через метод forward() класса RNN, где каждый шаг времени хранится в массиве states. Это необходимо для последующего обратного распространения ошибок.

Обратное распространение (Backward Pass)

Обратное распространение через время (BPTT — Backpropagation Through Time) является ключевым элементом обучения Vanilla RNN. Основные этапы:

  1. Инициализация градиентов скрытых состояний на последнем шаге dh_next = 0.

  2. Проход по последовательности в обратном порядке:

    [ dh = dh_{} + dh_{}] [ dh_{} = (1 - h_t^2) dh] [ dWx += dh_{} x_t^T] [ dWh += dh_{} h_{t-1}^T] [ db += dh_{}] [ dh_{} = W_h^T dh_{}]

  3. Агрегирование градиентов по всей последовательности и обновление весов.

ConvNetJS автоматизирует эту процедуру в методе backward(), используя массив состояний, сохраненных во время прямого прохода, и массив градиентов выходного слоя.

Настройка и оптимизация

Для эффективного обучения RNN в ConvNetJS используются следующие параметры:

  • learning_rate — скорость обучения, обычно малые значения, чтобы избежать нестабильности градиентов.
  • momentum — ускорение сходимости путем сглаживания обновлений.
  • clipval — ограничение градиентов, предотвращающее их взрыв. В Vanilla RNN критически важно, так как большие градиенты на длинных последовательностях вызывают неустойчивость.

Пример инициализации сети:

var rnn = new convnetjs.RNN({
  input_size: vocab_size,
  hidden_size: 128,
  output_size: vocab_size,
  learning_rate: 0.01,
  clipval: 5.0
});

Работа с последовательностями

RNN в ConvNetJS обрабатывает последовательности фиксированной длины. Каждое состояние скрытого слоя передается на следующий шаг. Для генерации текста или прогнозирования временных рядов последовательность подается пошагово, а выход предыдущего шага можно использовать как вход на следующий.

Хранение и повторное использование состояний скрытого слоя позволяет:

  • Генерировать последовательности произвольной длины.
  • Поддерживать контекст при работе с потоковыми данными.
  • Реализовать обучение с усеченным BPTT, когда градиенты рассчитываются только по ограниченному числу шагов.

Применение Vanilla RNN

Vanilla RNN применяется для задач, где важен порядок входных данных:

  • Генерация текста по обученному корпусу.
  • Анализ временных рядов, прогнозирование и детекция аномалий.
  • Простые задачи классификации последовательностей, например, распознавание рукописных цифр, представленных как последовательность координат.

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

Пример использования

var rnn = new convnetjs.RNN({input_size: 10, hidden_size: 50, output_size: 10});
var x = new convnetjs.Vol([/* входные данные */]);
rnn.forward(x);
rnn.backward();

Эта простая схема демонстрирует цикл прямого и обратного прохода для одного шага времени. Для последовательностей необходимо обернуть её в цикл по всем временным шагам, сохраняя состояния скрытого слоя для корректного BPTT.

Vanilla RNN в ConvNetJS сочетает наглядность реализации с гибкостью экспериментов, позволяя изучать динамику рекуррентных сетей и их обучение без сложных зависимостей и внешних библиотек.