train — основной метод библиотеки Brain.js для обучения
нейронной сети на предоставленных данных. Он принимает на вход массив
объектов с полями input и output, а также
опциональные настройки, позволяющие управлять процессом обучения и
точностью модели.
Для корректного обучения требуется структурировать данные следующим образом:
const trainingData = [
{ input: [0, 0], output: [0] },
{ input: [0, 1], output: [1] },
{ input: [1, 0], output: [1] },
{ input: [1, 1], output: [0] }
];
input — массив или объект с признаками.output — массив или объект с ожидаемым
результатом.trainМетод train имеет второй аргумент — объект опций,
который позволяет контролировать процесс обучения:
network.train(trainingData, {
iterations: 20000,
errorThresh: 0.005,
log: true,
logPeriod: 100,
learningRate: 0.3,
momentum: 0.1,
callback: null,
callbackPeriod: 10,
timeout: Infinity
});
Разберем каждый параметр подробно:
iterations — максимальное
количество итераций обучения. Если задано слишком мало, сеть может не
достичь нужной точности, слишком большое значение увеличивает время
обучения.
errorThresh — порог ошибки, при
достижении которого обучение считается завершённым. Величина обычно
задается в диапазоне от 0 до 1; чем меньше значение, тем точнее сеть, но
тем дольше обучение.
log — позволяет включить вывод
прогресса обучения. Может быть булевым значением или функцией для
кастомного логирования.
logPeriod — интервал итераций,
через который выводится лог, если log включен. Полезно для
отслеживания прогресса на больших данных.
learningRate — коэффициент
обучения, определяющий скорость обновления весов сети. Обычно находится
в диапазоне 0.1–0.5, но может подбираться экспериментально.
momentum — дополнительный
коэффициент, смягчающий изменения весов между итерациями. Позволяет
ускорить обучение и уменьшить вероятность застревания в локальных
минимумах.
callback — функция, вызываемая
через каждые callbackPeriod итераций, передающая текущую
ошибку сети. Может использоваться для динамического контроля
обучения.
callbackPeriod — количество
итераций между вызовами callback.
timeout — максимальное время
обучения в миллисекундах. Если время превышено, обучение
останавливается.
trainБалансировка данных Для корректного обучения желательно, чтобы набор данных был сбалансирован по классам. Несбалансированные данные могут привести к смещению сети в сторону наиболее часто встречающихся значений.
Нормализация входных данных Нейронные сети Brain.js лучше обучаются на данных с диапазоном значений от 0 до 1. Это особенно важно для числовых признаков, чтобы избежать проблем с градиентами.
Выбор функции активации В зависимости от типа
данных можно выбрать sigmoid, relu или
tanh. Для бинарных задач чаще используется
sigmoid, для многоклассовых — softmax (в
расширенных реализациях).
Отслеживание ошибки Метод train
возвращает объект с результатами обучения, включая достигнутую ошибку.
Это позволяет оценить, насколько хорошо сеть обучилась на
предоставленных данных.
const brain = require('brain.js');
const network = new brain.NeuralNetwork();
const trainingData = [
{ input: { a: 0, b: 0 }, output: { c: 0 } },
{ input: { a: 0, b: 1 }, output: { c: 1 } },
{ input: { a: 1, b: 0 }, output: { c: 1 } },
{ input: { a: 1, b: 1 }, output: { c: 0 } }
];
network.train(trainingData, {
iterations: 10000,
errorThresh: 0.002,
log: (error) => console.log('Ошибка сети:', error),
logPeriod: 200,
learningRate: 0.4,
momentum: 0.2,
callbackPeriod: 100
});
const output = network.run({ a: 1, b: 0 });
console.log('Результат предсказания:', output);
В этом примере демонстрируется:
train.learningRate
увеличивают точность, но замедляют обучение.iterations полезны для сложных
сетей, но при избытке может возникнуть переобучение.momentum помогает
ускорить сходимость и стабилизировать процесс обновления весов.Метод train является гибким инструментом для обучения
нейронных сетей в Brain.js, обеспечивая контроль над точностью,
скоростью и стабильностью процесса. Подбор оптимальных параметров
зависит от структуры сети, сложности задачи и объёма данных.