Построение REST API с предсказаниями

Brain.js — это библиотека для создания нейронных сетей на JavaScript. Она позволяет реализовывать различные типы сетей, включая полносвязные (Feedforward) и рекуррентные (Recurrent), для задач классификации, регрессии и предсказаний. В контексте REST API библиотека обеспечивает возможность обработки входных данных и выдачи прогнозов в реальном времени.


Выбор типа сети

Для большинства задач предсказаний чаще всего используют следующие виды сетей:

  • feedforward (NeuralNetwork): подходит для задач классификации и регрессии, когда данные не имеют временной зависимости.
  • recurrent (recurrent.LSTM или recurrent.GRU): применяется для работы с последовательностями, временными рядами, текстовыми данными.
const brain = require('brain.js');

// Пример feedforward сети
const net = new brain.NeuralNetwork({
  hiddenLayers: [10, 10], // две скрытые слоя по 10 нейронов
  activation: 'relu' // функция активации
});

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

  • hiddenLayers — массив, определяющий количество нейронов в каждом скрытом слое.
  • activation — функция активации, может быть sigmoid, relu или leaky-relu.
  • learningRate — скорость обучения сети, важный параметр для convergence при тренировке.

Подготовка данных

Brain.js требует определённого формата данных:

  • Для классификации входные данные должны быть числовыми, а выход — либо числами, либо объектом с вероятностями классов.
  • Для регрессии и предсказания временных рядов входные и выходные данные представляются как массивы чисел.

Пример подготовки данных для классификации:

const trainingData = [
  { input: { red: 1, green: 0, blue: 0 }, output: { color: 'red' } },
  { input: { red: 0, green: 1, blue: 0 }, output: { color: 'green' } },
  { input: { red: 0, green: 0, blue: 1 }, output: { color: 'blue' } }
];

Для временных рядов:

const timeSeriesData = [
  { input: [1, 2, 3], output: [4] },
  { input: [2, 3, 4], output: [5] }
];

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

Обучение производится методом train или асинхронным методом trainAsync:

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

Параметры:

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

Интеграция с REST API

REST API создается с использованием фреймворка Express. Сеть Brain.js можно встроить в обработчики маршрутов для обработки POST-запросов.

const express = require('express');
const bodyParser = require('body-parser');

const app = express();
app.use(bodyParser.json());

app.post('/predict', (req, res) => {
  const inputData = req.body;
  const output = net.run(inputData);
  res.json({ prediction: output });
});

app.listen(3000, () => {
  console.log('API запущено на порту 3000');
});

Особенности:

  • bodyParser.json() обеспечивает корректный парсинг JSON-запросов.
  • req.body содержит данные, переданные клиентом.
  • Метод net.run(input) возвращает предсказание, готовое для отправки в ответе API.

Работа с сохранением и загрузкой модели

Brain.js позволяет сохранять и загружать обученные модели, что особенно важно для REST API, чтобы не обучать сеть при каждом запуске.

// Сохранение модели
const jsonModel = net.toJSON();
require('fs').writeFileSync('model.json', JSON.stringify(jsonModel));

// Загрузка модели
const savedModel = JSON.parse(require('fs').readFileSync('model.json'));
net.fromJSON(savedModel);

Эти операции обеспечивают:

  • Быстрый старт API без повторного обучения.
  • Возможность обновления модели путем периодического обучения на новых данных.

Масштабирование и оптимизация

При использовании Brain.js в серверных приложениях важно учитывать:

  • Ограничения по объему входных данных: сети большого размера требуют значительных ресурсов.
  • Асинхронное обучение с использованием trainAsync для предотвращения блокировки event loop.
  • Кэширование предсказаний при повторяющихся запросах для ускорения работы API.

Пример расширенной интеграции с REST API

app.post('/predict-batch', async (req, res) => {
  const batch = req.body.inputs; // массив объектов input
  const predictions = batch.map(input => net.run(input));
  res.json({ predictions });
});

Особенности подхода:

  • Позволяет обрабатывать массивы запросов за один POST.
  • Снижает нагрузку на сервер при множественных последовательных запросах.

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

Для предсказания последовательностей можно использовать recurrent.LSTM:

const lstm = new brain.recurrent.LSTMTimeStep();

lstm.train([
  [1, 2, 3, 4],
  [2, 3, 4, 5],
  [3, 4, 5, 6]
], {
  learningRate: 0.01,
  iterations: 5000
});

const future = lstm.run([4, 5, 6]); // предсказание следующего значения

Ключевые моменты:

  • LSTMTimeStep удобен для прогнозирования числовых последовательностей.
  • Метод run возвращает следующую точку или массив точек, в зависимости от задачи.
  • В API можно обрабатывать запросы с временными рядами аналогично обычным входным данным.

Обработка ошибок и валидация данных

Для стабильной работы REST API важно проверять входные данные перед передачей их в сеть:

app.post('/predict', (req, res) => {
  const inputData = req.body;
  if (!inputData || typeof inputData !== 'object') {
    return res.status(400).json({ error: 'Некорректные данные' });
  }
  const output = net.run(inputData);
  res.json({ prediction: output });
});

Такой подход предотвращает падение сервера и обеспечивает корректные ответы клиенту.


Этот подход позволяет строить производительные REST API с предсказаниями на основе Brain.js, обеспечивая гибкость в выборе сетей, обработке данных и масштабировании приложений.