Кривые обучения: loss и метрики

Кривые обучения отображают процесс обучения модели в виде изменения значений функции потерь (loss) и метрик на обучающей и проверочной выборках. Эти кривые позволяют визуализировать качество обучения, выявлять переобучение или недообучение и настраивать гиперпараметры модели. В TensorFlow.js работа с кривыми обучения тесно связана с использованием методов fit и колбэков (callbacks).


Функция потерь (Loss)

Функция потерь количественно оценивает расхождение между предсказанными значениями модели и истинными метками. Выбор функции потерь зависит от задачи:

  • Регрессия: tf.losses.meanSquaredError, tf.losses.meanAbsoluteError
  • Классификация: tf.losses.softmaxCrossEntropy, tf.losses.sigmoidCrossEntropy

В TensorFlow.js при компиляции модели функция потерь задается через метод compile:

model.compile({
  optimizer: 'adam',
  loss: 'categoricalCrossentropy',
  metrics: ['accuracy']
});

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


Метрики (Metrics)

Метрики позволяют контролировать качество модели на разных стадиях обучения. Основная метрика для классификации — точность (accuracy), для регрессии — средняя абсолютная ошибка (mae) или среднеквадратичная ошибка (mse). Метрики задаются аналогично функции потерь:

metrics: ['accuracy', 'mse']

В отличие от функции потерь, метрики не участвуют напрямую в процессе оптимизации, но дают представление о практическом качестве модели.


История обучения и кривые

Метод model.fit возвращает объект History, содержащий значения loss и метрик по эпохам:

const history = await model.fit(xTrain, yTrain, {
  epochs: 50,
  validationData: [xVal, yVal],
  callbacks: tf.callbacks.earlyStopping({monitor: 'val_loss'})
});

Объект history.history имеет вид:

{
  loss: [0.9, 0.8, 0.7, ...],
  accuracy: [0.5, 0.6, 0.65, ...],
  val_loss: [1.0, 0.85, 0.75, ...],
  val_accuracy: [0.45, 0.55, 0.6, ...]
}

Эти массивы можно использовать для построения графиков кривых обучения с помощью библиотек визуализации, например, Plotly.js или Chart.js.


Анализ кривых обучения

Ключевые признаки:

  • Недообучение (Underfitting): высокий loss на обучающей выборке, низкая точность; кривые loss и метрики обучающей и проверочной выборки сходны и не улучшаются.
  • Переобучение (Overfitting): loss на обучающей выборке снижается, а на проверочной стабилизируется или растет; точность на обучающей выборке выше, чем на проверочной.
  • Оптимальное обучение: loss на обучающей и проверочной выборках сходятся к минимальным значениям, метрики стабилизируются на высоком уровне.

Для корректного анализа необходимо следить за масштабом графиков и количеством эпох. Часто имеет смысл применять регуляризацию (dropout, l2) или методы ранней остановки (EarlyStopping) для предотвращения переобучения.


Визуализация кривых в TensorFlow.js

Для интерактивного анализа используют JavaScript-библиотеки визуализации:

const epochs = history.epoch;
const loss = history.history.loss;
const valLoss = history.history.val_loss;

const trace1 = { x: epochs, y: loss, type: 'scatter', name: 'train_loss' };
const trace2 = { x: epochs, y: valLoss, type: 'scatter', name: 'val_loss' };

Plotly.newPlot('plotDiv', [trace1, trace2], {title: 'Кривые обучения'});

График позволяет быстро увидеть, на какой эпохе происходит стабилизация или резкий рост потерь.


Роль колбэков (Callbacks)

Колбэки помогают автоматически управлять обучением:

  • tf.callbacks.earlyStopping — прекращает обучение при отсутствии улучшений на проверочной выборке.
  • tf.callbacks.modelCheckpoint — сохраняет модель при улучшении метрики.
  • tf.callbacks.tensorBoard — интеграция с TensorBoard для визуального анализа кривых в браузере.

Пример ранней остановки:

const earlyStopping = tf.callbacks.earlyStopping({
  monitor: 'val_loss',
  patience: 5
});

await model.fit(xTrain, yTrain, {
  epochs: 100,
  validationData: [xVal, yVal],
  callbacks: [earlyStopping]
});

Практические рекомендации

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

Кривые обучения являются основным инструментом контроля качества и стабильности модели в TensorFlow.js. Их анализ позволяет оптимизировать архитектуру сети, гиперпараметры и процесс обучения.