Регрессия на новых примерах: predict

Библиотека ml5.js предоставляет удобные высокоуровневые инструменты для машинного обучения в браузере, упрощая работу с нейронными сетями для задач классификации, регрессии, генерации изображений и звуков. В контексте регрессии основным объектом является ml5.neuralNetwork, который позволяет строить модели, обучать их и получать предсказания на новых данных.

Создание модели регрессии

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

const options = {
  task: 'regression',
  debug: true
};

const nn = ml5.neuralNetwork(options);

Здесь task: 'regression' определяет, что модель будет предсказывать не категории, а непрерывные численные значения. Параметр debug: true включает подробный вывод в консоль, что удобно для отслеживания процесса обучения.

Добавление данных

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

nn.addData({ x: 1, y: 2 }, { z: 3 });
nn.addData({ x: 2, y: 4 }, { z: 6 });
nn.addData({ x: 3, y: 6 }, { z: 9 });

Каждый объект входных данных { x, y } соответствует объекту целевых значений { z }. Важный момент: структура данных для обучения должна совпадать со структурой данных, которые будут использоваться для предсказания.

Нормализация данных

Перед обучением данные обычно нормализуются:

nn.normalizeData();

Нормализация улучшает сходимость обучения и повышает точность предсказаний. Ml5.js поддерживает автоматическую нормализацию через этот метод.

Обучение модели

Обучение происходит с использованием метода train, который принимает параметры оптимизации:

const trainingOptions = {
  epochs: 50,
  batchSize: 12
};

nn.train(trainingOptions, finishedTraining);

function finishedTraining() {
  console.log('Обучение завершено');
}
  • epochs — количество полных проходов по всем обучающим данным.
  • batchSize — размер мини-батча, на котором обновляются веса сети.

Предсказание новых значений: predict

Метод predict используется для вычисления результата на новых данных, которые не были включены в обучающую выборку. Синтаксис:

nn.predict({ x: 4, y: 8 }, gotResult);

function gotResult(error, results) {
  if (error) {
    console.error(error);
    return;
  }
  console.log(results);
}

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

  • Структура входных данных должна точно соответствовать структуре, использованной при обучении ({ x, y } в примере).
  • Результат возвращается массивом объектов. Для задачи регрессии каждый объект содержит предсказанные значения целевых переменных, например results[0].value.

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

[ { z: 12 } ]

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

predict поддерживает передачу массива объектов для пакетного предсказания:

const inputs = [
  { x: 5, y: 10 },
  { x: 6, y: 12 }
];

nn.predict(inputs, (err, results) => {
  if (err) return console.error(err);
  console.log(results);
});

Возвращается массив предсказаний, где каждый элемент соответствует объекту входных данных по позиции в массиве.

Использование предсказаний в визуализации

Одним из ключевых преимуществ ml5.js является возможность использовать результаты предсказаний в интерактивных веб-приложениях, например, для графиков или управления элементами интерфейса:

nn.predict({ x: mouseX, y: mouseY }, (err, results) => {
  if (!err) {
    const zPred = results[0].z;
    ellipse(mouseX, mouseY, zPred, zPred);
  }
});

В этом примере значение zPred, предсказанное моделью, используется для динамического изменения размера графического объекта на холсте.

Синхронное использование predict с промисами

Модель можно использовать с асинхронной записью, что упрощает обработку предсказаний:

async function makePrediction(input) {
  const results = await nn.predict(input);
  console.log(results);
}

makePrediction({ x: 7, y: 14 });

Асинхронный подход упрощает интеграцию модели в современные веб-приложения с использованием async/await.

Особенности точности предсказаний

Точность модели напрямую зависит от:

  • качества и объёма обучающих данных,
  • количества эпох обучения,
  • структуры нейронной сети (количество скрытых слоёв и нейронов),
  • нормализации данных.

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

Итоговое применение

Метод predict является ключевым инструментом для реального использования обученной модели. Он позволяет интегрировать регрессионные модели в визуализации, приложения реального времени и интерактивные интерфейсы, обеспечивая динамическое взаимодействие между данными и пользователем.

Использование predict в ml5.js отличается простотой, гибкостью и полной совместимостью с веб-экосистемой JavaScript, что делает библиотеку удобной платформой для обучения и внедрения регрессионных моделей.