Класс Trainer

Класс Trainer является важным компонентом библиотеки ConvNetJS, отвечающим за обучение нейронных сетей. Он инкапсулирует процесс оптимизации весов сети, управление скоростью обучения, методами регуляризации и параметрами обновления градиентов. В отличие от ручного обновления весов через вызовы forward и backward, Trainer автоматизирует этот процесс, предоставляя удобный интерфейс для настройки и контроля тренировки.


Инициализация Trainer

Создание экземпляра Trainer происходит через конструктор, который принимает два ключевых аргумента: объект нейронной сети (net) и словарь с параметрами (options):

var trainer = new convnetjs.Trainer(net, {
  method: 'sgd',       // метод оптимизации
  learning_rate: 0.01, // скорость обучения
  momentum: 0.9,       // моментум для SGD
  batch_size: 1,       // размер батча
  l2_decay: 0.0001     // коэффициент L2-регуляризации
});

Основные параметры конструктора:

  • method — метод оптимизации. Поддерживаются:

    • 'sgd' — стохастический градиентный спуск.
    • 'adagrad' — адаптивный градиентный спуск.
    • 'adam' — адаптивный метод с моментумом и нормализацией.
    • 'windowgrad' — усреднённый градиент по окну.
  • learning_rate — скорость обучения, контролирующая размер шага при обновлении весов.

  • momentum — коэффициент моментума, ускоряющий сходимость в направлениях с устойчивыми градиентами.

  • batch_size — количество образцов, обрабатываемых за один шаг обновления.

  • l1_decay, l2_decay — коэффициенты регуляризации для предотвращения переобучения.

  • clip_gradients — предел для отсечения слишком больших градиентов.


Основные методы

train(x, y)

Метод train выполняет полный цикл обучения на одном батче входных данных x с правильными метками y.

  • x — входные данные. В случае классификации обычно одномерный массив или Vol.
  • y — правильные метки или целевые значения.

Возвращает объект, содержащий:

  • loss — значение функции потерь для данного батча.
  • gradients — градиенты по весам (опционально, чаще используется для отладки).

Пример использования:

var stats = trainer.train(x, y);
console.log('Loss: ', stats.loss);
trainBatch(x_batch, y_batch)

Позволяет тренировать сеть на множестве примеров одновременно, ускоряя обучение при использовании batch_size > 1.

  • x_batch — массив входных векторов или Vol объектов.
  • y_batch — массив целевых меток.

Метод автоматически усредняет градиенты по батчу перед обновлением весов.

trainForwardBackward(x, y)

Интерфейс низкого уровня для ручного управления обучением. Выполняет:

  1. Прямое распространение (forward) по сети.
  2. Обратное распространение (backward) с вычислением градиентов.
  3. Обновление весов с использованием текущего метода оптимизации.

Методы оптимизации

Trainer инкапсулирует несколько алгоритмов обновления весов:

  • SGD (Stochastic Gradient Descent) Обновление весов по формуле: [ w = w - g + w_{}] где () — learning_rate, (g) — градиент, () — momentum.

  • AdaGrad Поддерживает адаптивную скорость обучения для каждого веса, уменьшая шаги для часто обновляемых параметров.

  • Adam Комбинирует моментум и адаптивное масштабирование градиентов, что повышает стабильность при обучении глубоких сетей.

  • WindowGrad Учитывает среднее значение градиентов по окну, сглаживая шум.


Регуляризация и усечение градиентов

Для предотвращения переобучения Trainer поддерживает L1 и L2 регуляризацию, а также усечение градиентов:

var trainer = new convnetjs.Trainer(net, {
  l1_decay: 0.001,
  l2_decay: 0.0001,
  clip_gradients: 5.0
});
  • L1/L2 decay — добавляют штраф к функции потерь пропорционально весам, уменьшая их значения.
  • clip_gradients — ограничивает максимальное значение градиента, предотвращая взрыв градиентов.

Сбор статистики обучения

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

  • loss — текущее значение функции потерь.
  • gradients — градиенты весов (для анализа или визуализации).
  • iteration — номер текущей итерации обучения.
  • batch_size — размер батча для данного шага.

Пример вывода статистики:

for (var i = 0; i < 1000; i++) {
    var stats = trainer.train(x, y);
    if (i % 100 === 0) {
        console.log('Iteration', i, 'Loss', stats.loss);
    }
}

Интеграция с Vol и Net

Trainer тесно связан с классами Vol (структура для хранения данных и градиентов) и Net (нейронная сеть). Его методы используют эти объекты для передачи данных и градиентов. Внутри каждого шага обучения:

  1. Входной Vol передаётся в net.forward().
  2. Результат сравнивается с правильной меткой y.
  3. Вычисляется градиент функции потерь.
  4. Вызывается net.backward() для распространения градиентов.
  5. Обновляются веса с учётом выбранного метода оптимизации.

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

  • Метод оптимизации выбирается в зависимости от сложности сети и данных: для небольших сетей подходит 'sgd', для глубоких 'adam'.
  • learning_rate — ключевой параметр. Малое значение замедляет обучение, большое может вызвать расходимость.
  • momentum помогает преодолевать локальные минимумы и ускоряет обучение.
  • batch_size балансирует скорость и качество: маленький батч даёт шумные градиенты, большой — медленнее обновления.
  • Регуляризация особенно важна для малых данных или глубоких сетей.

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