Развертка рекуррентной сети

ConvNetJS — это чисто JavaScript-библиотека для построения и обучения нейронных сетей, включая полносвязные, сверточные и рекуррентные сети. Основной строительный блок — объект Net, который хранит слои и веса, а также обеспечивает интерфейс для прямого и обратного распространения.

Рекуррентные сети (RNN) в ConvNetJS строятся на базе слоя LSTM или RNN, что позволяет моделировать последовательности данных. Каждый рекуррентный слой состоит из ячейки состояния, которая хранит внутреннюю память, и входного, выходного и скрытого слоев, соединённых весами. Сеть способна учитывать прошлую информацию при обработке текущего элемента последовательности.

Создание рекуррентной сети

var net = new convnetjs.Net();
var layer_defs = [];
layer_defs.push({type:'input', out_sx:1, out_sy:1, out_depth:1});
layer_defs.push({type:'lstm', num_units:20});
layer_defs.push({type:'regression', num_neurons:1});
net.makeLayers(layer_defs);
  • input — входной слой, определяет размер входного вектора.
  • lstm — рекуррентный слой с LSTM-ячейками, где num_units — число скрытых нейронов.
  • regression — выходной слой для предсказания числовых значений.

Обучение на последовательностях

Обучение рекуррентной сети требует подачи данных в виде последовательностей. ConvNetJS использует метод forward() для прямого распространения и backward() для обратного распространения ошибки через время (BPTT — Backpropagation Through Time).

var trainer = new convnetjs.SGDTrainer(net, {learning_rate:0.01, momentum:0.9, batch_size:1, l2_decay:0.001});

for(var t=0;t<sequence.length;t++){
    var x = new convnetjs.Vol([sequence[t]]);
    net.forward(x);
    var loss = trainer.train(x, target[t]);
}
  • sequence — массив входных данных.
  • target — массив целевых значений.
  • trainer.train() автоматически вызывает прямое и обратное распространение для обновления весов.

Управление памятью LSTM

LSTM-слой хранит состояния ячеек (cell state) и скрытые состояния (hidden state). ConvNetJS позволяет сбрасывать состояние между последовательностями, чтобы избежать «загрязнения» данных из предыдущих примеров:

net.layers[1].reset_state();

Важно использовать эту функцию при обучении на независимых последовательностях.

Генерация последовательностей

После обучения сеть можно использовать для предсказания последовательностей. Метод forward() применяется поэлементно:

var x = new convnetjs.Vol([start_value]);
for(var t=0; t<50; t++){
    net.forward(x);
    var prediction = net.layers[2].out_act.w[0];
    x = new convnetjs.Vol([prediction]);
}

Таким образом можно генерировать новые последовательности на основе исходного значения или предыдущего предсказания.

Настройка гиперпараметров

Ключевые гиперпараметры для рекуррентной сети в ConvNetJS:

  • num_units — количество нейронов в LSTM; большее значение повышает способность запоминания, но увеличивает риск переобучения.
  • learning_rate — скорость обучения; рекомендуется начинать с малых значений (0.01–0.001) для стабильного градиентного спуска.
  • momentum — ускоряет сходимость, учитывая предыдущие шаги градиента.
  • l2_decay — регуляризация весов для предотвращения переобучения.
  • batch_size — для RNN обычно выбирают 1, чтобы обучение проходило по шагам последовательности.

Важные особенности работы RNN в ConvNetJS

  • Обратное распространение ошибки выполняется через все шаги последовательности; слишком длинные последовательности могут привести к затухающим или взрывающимся градиентам.
  • Для стабильного обучения рекомендуется нормализовать входные данные и иногда использовать усечение градиентов (gradient clipping), реализуемое вручную через изменение весов после backward().
  • LSTM лучше справляется с долгосрочными зависимостями, чем стандартный RNN-слой, благодаря контролю памяти через входные, забывающие и выходные гейты.

Визуализация и отладка

ConvNetJS предоставляет встроенные средства визуализации:

var vis = new convnetjs.Visualizer(net);
vis.drawGraph();

Можно отслеживать:

  • изменения весов,
  • активации слоев,
  • кривую потерь (loss) во времени.

Это упрощает выявление проблем с переобучением или сходимостью.


Если требуется, можно дополнительно раскрыть пример работы с текстовыми последовательностями или временными рядами, включая предобработку данных и преобразование их в Vol объекты для подачи в сеть.

Хотите, чтобы я включил этот практический пример?