Обучение рекуррентных сетей в Brain.js

Рекуррентные нейронные сети (RNN, Recurrent Neural Networks) предназначены для работы с последовательными данными, где значение текущего шага зависит от предыдущих. В Brain.js это реализуется через класс recurrent.LSTM и его варианты, позволяя строить модели для анализа текста, временных рядов, прогнозирования и генерации последовательностей.

Создание сети

Для начала создаётся объект сети:

const brain = require('brain.js');
const net = new brain.recurrent.LSTM();

Ключевые параметры:

  • inputSize — размер входного вектора (если вход нормализован, обычно не указывают).
  • hiddenLayers — массив чисел, задающих количество нейронов в скрытых слоях.
  • outputSize — размер выходного вектора.
  • learningRate — скорость обучения, по умолчанию 0.005.

Пример с указанием скрытых слоёв:

const net = new brain.recurrent.LSTM({
  hiddenLayers: [20, 20],
  learningRate: 0.01
});

Формат данных для обучения

RNN в Brain.js может работать с последовательностями текста или числовыми последовательностями. Для текстовых данных используется массив объектов вида:

[
  { input: "Привет", output: "Здравствуйте" },
  { input: "Пока", output: "До свидания" }
]

Для числовых последовательностей:

[
  { input: [0, 1, 2], output: [3] },
  { input: [1, 2, 3], output: [4] }
]

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

Обучение сети

Обучение выполняется методом train:

net.train(data, {
  iterations: 2000,
  log: true,
  logPeriod: 100,
  errorThresh: 0.005
});

Основные параметры:

  • iterations — максимальное количество эпох обучения.
  • log — вывод прогресса в консоль.
  • logPeriod — частота логирования.
  • errorThresh — порог ошибки для остановки обучения.
  • learningRate — скорость обучения.

Brain.js использует алгоритм обратного распространения через время (BPTT, Backpropagation Through Time) для корректировки весов сети, учитывая зависимость последовательностей.

Применение обученной сети

После обучения сеть готова к предсказаниям с использованием метода run:

const output = net.run("Привет");
console.log(output); // "Здравствуйте"

Для генерации числовых последовательностей:

const output = net.run([2, 3, 4]);
console.log(output); // [5]

Настройка рекуррентной сети

Размер скрытых слоёв: от 10 до 100 нейронов на слой оптимален для небольших задач. Слишком маленький размер ведёт к недообучению, слишком большой — к переобучению и увеличенному времени обучения.

Количество слоёв: чаще всего используется 1–3 слоя. Дополнительные слои повышают выразительность модели, но увеличивают сложность и требования к данным.

Длина последовательности: LSTM лучше работает с ограниченной длиной входной последовательности. Для длинных текстов рекомендуется разбивать их на фрагменты.

Преобразование текста в формат сети

Brain.js автоматически обрабатывает строки как последовательность символов. Для более сложного контроля можно использовать токенизацию или векторизацию:

const tokenized = text.split('').map(char => char.charCodeAt(0) / 255);

Это позволяет нормализовать входы в диапазон [0,1] и улучшить стабильность обучения.

Сохранение и загрузка модели

Сеть можно сериализовать в JSON для хранения:

const json = net.toJSON();

И восстановить позже:

const net2 = new brain.recurrent.LSTM();
net2.fromJSON(json);

Это удобно для переноса модели между проектами или для использования на сервере и клиенте.

Продвинутые возможности

LSTMTimeStep — специализированный класс для работы с временными рядами числовых данных. Он автоматически учитывает последовательность шагов времени, облегчая предсказания на основе истории:

const timeNet = new brain.recurrent.LSTMTimeStep();
timeNet.train([
  [1,2,3],
  [2,3,4],
  [3,4,5]
]);
const prediction = timeNet.run([4,5]);

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

Регуляризация: Brain.js поддерживает dropout через параметр decayRate для уменьшения переобучения.

Рекомендации по эффективности

  • Использовать нормализацию данных для числовых последовательностей.
  • Разбивать длинные тексты на логические сегменты.
  • Контролировать размер и количество скрытых слоёв для оптимального баланса скорости и точности.
  • Использовать логирование процесса обучения для мониторинга ошибок и корректировки параметров.

Рекуррентные сети в Brain.js предоставляют простой и понятный интерфейс для решения задач с последовательными данными, обеспечивая возможность работать как с текстом, так и с числовыми рядами. Их гибкость позволяет создавать предсказательные модели без необходимости глубокого погружения в низкоуровневую математику LSTM.