Посимвольные модели на LSTM

LSTM (Long Short-Term Memory) — это разновидность рекуррентных нейронных сетей (RNN), предназначенная для работы с последовательными данными, такими как текст. В отличие от стандартных RNN, LSTM способны эффективно запоминать длинные зависимости в последовательностях, что делает их идеальными для задач генерации текста посимвольно.

Brain.js предоставляет возможность создавать LSTM-сети для работы с текстом через объект recurrent.LSTM. Эти сети обучаются на последовательностях символов, а затем могут генерировать новые последовательности на основе изученного контекста.


Создание и конфигурация сети

Для начала необходимо импортировать библиотеку и создать объект LSTM:

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

Основные параметры, которые можно настроить при создании LSTM:

  • inputSize — размер входного слоя (опционально; Brain.js обычно определяет автоматически).
  • hiddenLayers — массив, задающий количество и размер скрытых слоев.
  • outputSize — размер выходного слоя (для посимвольной модели равен размеру алфавита).
  • learningRate — скорость обучения (по умолчанию 0.01).
  • decayRate — коэффициент затухания ошибки для стабилизации обучения.

Пример с пользовательскими скрытыми слоями:

const lstm = new brain.recurrent.LSTM({
  hiddenLayers: [128, 128], // два скрытых слоя по 128 нейронов
  learningRate: 0.005
});

Подготовка данных для посимвольного обучения

Для LSTM важно, чтобы данные представляли собой последовательности символов. В Brain.js это достигается через массив объектов с полями input и output:

const trainingData = [
  { input: "привет", output: "привет" },
  { input: "мир", output: "мир" },
  { input: "hello", output: "hello" }
];

Особенности подготовки:

  • Сохранять последовательность символов: каждый input должен содержать последовательность, которую сеть должна научиться предсказывать.
  • Выход совпадает с входом для генерации: в посимвольной генерации сеть обучается воспроизводить последовательность и предсказывать следующий символ.
  • Нормализация текста: можно привести все символы к нижнему регистру, удалить лишние символы, чтобы уменьшить размер алфавита.

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

Обучение LSTM выполняется методом train, принимающим массив данных и объект конфигурации обучения:

lstm.train(trainingData, {
  iterations: 2000,
  log: true,
  logPeriod: 100,
  learningRate: 0.01,
  errorThresh: 0.005
});

Ключевые параметры обучения:

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

Примечание: посимвольные модели требуют больше итераций и меньшего learningRate, чем модели, обучающиеся на словах, так как алфавит может быть большим, а контекст сложным.


Генерация текста

После обучения сеть способна генерировать текст посимвольно. Метод run принимает начальную последовательность (seed) и возвращает предсказанную последовательность:

const output = lstm.run('прив');
console.log(output); // Например: "ет"

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

let seed = 'п';
let result = seed;

for (let i = 0; i < 50; i++) {
  const next = lstm.run(seed);
  result += next;
  seed = next; // следующий шаг строится на последнем символе
}

console.log(result);

Практические советы по улучшению качества

  • Увеличение объема данных: больше текстов приводит к более осмысленной генерации.
  • Увеличение скрытых слоев: сложные последовательности требуют более глубоких сетей.
  • Регуляризация: небольшая скорость обучения и более длинные эпохи помогают избежать переобучения.
  • Разбиение на n-символьные последовательности: сеть легче учит короткие последовательности и лучше предсказывает следующий символ.

Хранение и восстановление модели

После обучения сеть можно сохранить в JSON и загрузить позже без повторного обучения:

const json = lstm.toJSON(); // экспорт модели
const lstm2 = new brain.recurrent.LSTM();
lstm2.fromJSON(json); // восстановление модели

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


Особенности работы с Brain.js

  • Автоматическая обработка алфавита: Brain.js автоматически индексирует символы для внутреннего представления.
  • Скорость работы: для больших моделей и длинных последовательностей рекомендуется использовать Node.js и отключать логирование.
  • Совместимость: LSTM Brain.js может использоваться как в серверных, так и в фронтенд-приложениях, но обучение больших моделей на клиенте неэффективно.

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

const trainingData = [
  { input: "мир", output: "мир" },
  { input: "машина", output: "машина" },
  { input: "марс", output: "марс" }
];

const lstm = new brain.recurrent.LSTM({
  hiddenLayers: [128, 128],
  learningRate: 0.005
});

lstm.train(trainingData, {
  iterations: 3000,
  log: true,
  logPeriod: 200,
  errorThresh: 0.005
});

const seed = 'м';
let text = seed;

for (let i = 0; i < 20; i++) {
  const next = lstm.run(seed);
  text += next;
  seed = next;
}

console.log(text); // Возможная генерация: "ирмашинамарс..."

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