Оценка точности модели

ConvNetJS — это чисто JavaScript библиотека для построения, обучения и оценки нейронных сетей, включая сверточные сети. Оценка точности модели является ключевым этапом в процессе разработки, поскольку позволяет количественно измерить способность сети корректно предсказывать классы на новых данных.


Метрики точности

В ConvNetJS основной метрикой оценки является accuracy — доля правильно классифицированных примеров. Для задачи классификации с (N) примерами и (C) классами точность вычисляется как:

[ = .]

Кроме точности, для углубленного анализа можно использовать:

  • Loss (потери) — среднее значение функции потерь по всем примерам. Для классификации обычно используется softmax loss.
  • Top-K accuracy — полезно в задачах с большим числом классов, измеряет вероятность того, что правильный класс находится среди K наиболее вероятных предсказаний.

Подготовка данных для оценки

Важный аспект — правильная подготовка тестового набора:

  1. Отделение тренировочных и тестовых данных. ConvNetJS не имеет встроенного разделения, поэтому нужно вручную формировать отдельные массивы данных для обучения и тестирования.
  2. Нормализация данных. Входные значения должны быть приведены к диапазону, ожидаемому сетью, например, ([0, 1]) или ([-1, 1]).
  3. Формирование мини-батчей. Для стабильной оценки больших наборов данных используется пакетная обработка. Размер батча влияет на скорость и точность вычислений.

Пример создания тестового набора:

var testData = [];
for (var i = 0; i < dataset.length; i++) {
    if (i % 5 === 0) { // каждый 5-й пример для теста
        testData.push(dataset[i]);
    }
}

Оценка точности модели

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

Пример вычисления точности:

var correct = 0;
for (var i = 0; i < testData.length; i++) {
    var x = testData[i].x;
    var y = testData[i].y;
    var predicted = net.forward(x);
    if (predicted.w.indexOf(Math.max(...predicted.w)) === y) {
        correct += 1;
    }
}
var accuracy = correct / testData.length;
console.log("Accuracy:", accuracy);

Пояснения к коду:

  • net.forward(x) возвращает объект Vol, содержащий предсказанные вероятности для каждого класса.
  • Использование Math.max(...predicted.w) позволяет определить класс с максимальной вероятностью.
  • Сравнение с истинным классом y позволяет подсчитать количество правильных предсказаний.

Визуализация потерь и точности

Для анализа поведения модели во время обучения удобно строить графики loss и accuracy по эпохам:

var stats = trainer.train(xBatch, yBatch);
lossHistory.push(stats.loss);
accuracyHistory.push(evaluateAccuracy(net, testData));

Графики позволяют выявить:

  • Переобучение (overfitting), если точность на тренировочном наборе растет, а на тестовом падает.
  • Недообучение (underfitting), если обе кривые низкие и не улучшаются.

Тонкости при оценке

  1. Shuffle данных. Тестовый набор должен быть случайным, чтобы исключить систематические ошибки.
  2. Многократная проверка. Для точной оценки рекомендуется повторять тест несколько раз с разными разбиениями данных.
  3. Мини-батчи и производительность. ConvNetJS позволяет вычислять forward-проход по батчу, что ускоряет оценку на больших данных:
for (var i = 0; i < testData.length; i += batchSize) {
    var batch = testData.slice(i, i + batchSize);
    batch.forEach(d => {
        var pred = net.forward(d.x);
        // подсчет правильных предсказаний
    });
}

Расширенные подходы

  • Confusion Matrix — матрица ошибок, полезна для анализа, какие классы сеть путает.
  • Precision и Recall — метрики для задач с несбалансированными классами.
  • ROC-AUC — для бинарной классификации с вероятностными предсказаниями.

В ConvNetJS для этих метрик приходится самостоятельно реализовывать вычисления на основе выходов сети и истинных меток, что позволяет гибко настраивать оценку под конкретную задачу.


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

  • Использовать отдельный тестовый набор, который не участвует в обучении.
  • Сохранять промежуточные веса сети для повторной оценки.
  • Сравнивать результаты по нескольким метрикам для комплексной оценки модели.
  • Отслеживать как тренировочные, так и тестовые потери, чтобы выявлять переобучение.

Тщательная оценка точности в ConvNetJS позволяет не только измерять текущие результаты сети, но и формировать стратегии дальнейшей оптимизации архитектуры и гиперпараметров.