Кросс-валидация — это метод оценки производительности модели машинного обучения, позволяющий эффективно использовать ограниченный набор данных. В браузерной среде с TensorFlow.js она обеспечивает возможность обучения и тестирования моделей непосредственно на стороне клиента без необходимости обращения к серверу.
Главная цель кросс-валидации — получить надежную оценку качества модели, минимизируя переобучение (overfitting) и недообучение (underfitting). На практике это достигается разбиением исходного набора данных на несколько непересекающихся подмножеств (folds), поочередным обучением модели на одних подмножествах и проверкой на других.
k-Fold Cross-Validation Данные делятся на (k) равных частей. Модель обучается (k) раз: каждый раз одна из частей используется как тестовая, а остальные (k-1) — как обучающие. Итоговая оценка производительности модели получается усреднением метрик по всем итерациям.
В TensorFlow.js это реализуется вручную, так как библиотека не предоставляет встроенной функции для автоматической кросс-валидации. Процесс включает:
tf.util.shuffle или аналогичных функций;k подмножеств;Leave-One-Out Cross-Validation (LOOCV) В этом подходе каждая точка данных последовательно используется как тестовая, а оставшиеся данные — как обучающие. Метод обеспечивает максимально точную оценку, но сильно нагружает браузер при больших наборах данных.
Stratified k-Fold Применяется для несбалансированных классов, когда важно сохранить пропорции классов в каждом подмножестве. В TensorFlow.js необходимо реализовать контроль распределения классов вручную при формировании фолдов.
Перед проведением кросс-валидации необходимо убедиться, что данные правильно подготовлены:
tf.tensor
и методы sub и div для вычитания среднего и
деления на стандартное отклонение.slice или
tf.gather позволяет выделять части тензора для обучения и
тестирования.Пример подхода для k-Fold Cross-Validation:
import * as tf from '@tensorflow/tfjs';
const k = 5;
const data = tf.tensor([...]); // исходный набор данных
const labels = tf.tensor([...]); // метки классов
// Перемешивание данных
const indices = tf.util.createShuffledIndices(data.shape[0]);
const foldSize = Math.floor(data.shape[0] / k);
let accuracies = [];
for (let i = 0; i < k; i++) {
const start = i * foldSize;
const end = start + foldSize;
const testIndices = indices.slice(start, end);
const trainIndices = indices.filter(idx => idx < start || idx >= end);
const xTrain = tf.gather(data, trainIndices);
const yTrain = tf.gather(labels, trainIndices);
const xTest = tf.gather(data, testIndices);
const yTest = tf.gather(labels, testIndices);
const model = createModel(); // функция создания модели
await model.fit(xTrain, yTrain, {epochs: 10, batchSize: 32});
const evalResult = model.evaluate(xTest, yTest);
accuracies.push(evalResult[1].dataSync()[0]); // точность
}
const meanAccuracy = accuracies.reduce((a, b) => a + b) / k;
console.log('Средняя точность модели:', meanAccuracy);
await model.fit() и tf.nextFrame() позволяют
не блокировать основной поток интерфейса.tfvis для мониторинга обучения модели на
каждом фолде.В k-Fold Cross-Validation результаты усредняются, что снижает влияние случайности:
Использование этих метрик в сочетании с кросс-валидацией дает более надежное представление о стабильности модели и позволяет выявить потенциальное переобучение.
fetch или
WebSockets для онлайн-обучения и тестирования модели на новых данных без
перезагрузки страницы.tf.data API позволяет создавать потоковые
датасеты и эффективно обрабатывать большие массивы данных.earlyStopping) помогает стабилизировать обучение при
ограниченных вычислительных ресурсах.Кросс-валидация в TensorFlow.js объединяет классические методы оценки модели с особенностями браузерной среды, обеспечивая гибкость, безопасность данных и возможность интерактивного анализа.