Обучение агента в ConvNetJS

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

Сеть в ConvNetJS строится с использованием объекта ConvNetJS.Net, который инициализируется с массивом слоёв. Каждый слой задаётся как объект с типом, параметрами фильтров, количеством нейронов и другими настройками:

var layer_defs = [];
layer_defs.push({type:'input', out_sx:28, out_sy:28, out_depth:1}); // размер входа
layer_defs.push({type:'conv', sx:5, filters:8, stride:1, pad:2, activation:'relu'});
layer_defs.push({type:'pool', sx:2, stride:2});
layer_defs.push({type:'softmax', num_classes:10});

var net = new convnetjs.Net();
net.makeLayers(layer_defs);

Ключевые моменты:

  • input задаёт размер входных данных.
  • conv выполняет свёртку с заданным числом фильтров, размером ядра и функцией активации.
  • pool уменьшает пространственные размеры карт признаков.
  • softmax служит для классификации на несколько классов.

Инициализация тренера и параметры обучения

Для обучения сети создаётся объект тренера ConvNetJS.SGDTrainer, где задаются метод оптимизации, скорость обучения, коэффициенты регуляризации и параметры момента:

var trainer = new convnetjs.SGDTrainer(net, {
  method: 'adadelta', 
  batch_size: 20, 
  l2_decay: 0.001
});

Важные параметры:

  • method: 'sgd', 'adadelta', 'adam' — выбор алгоритма оптимизации.
  • batch_size: количество образцов, обрабатываемых за один шаг обучения.
  • l2_decay: коэффициент L2-регуляризации для предотвращения переобучения.
  • momentum и learning_rate применяются при классическом SGD.

Подготовка данных

ConvNetJS работает с массивами чисел или объектами Vol, представляющими многомерные данные. Для обучения изображений часто используется преобразование данных в объекты Vol:

var x = new convnetjs.Vol(28, 28, 1);
for (var i=0;i<28*28;i++) {
  x.w[i] = pixels[i]/255.0; // нормализация
}

Особенности:

  • Все значения нормализуются в диапазон [0,1] или [-1,1] для стабильного обучения.
  • Vol хранит веса, градиенты и размерность данных.

Процесс обучения

Обучение выполняется методом train тренера, которому передаётся объект Vol и метка класса:

trainer.train(x, label);

Во время тренировки библиотека автоматически:

  • Вычисляет прямой проход (forward) по сети.
  • Рассчитывает функцию потерь (например, кросс-энтропию).
  • Производит обратный проход (backward) для обновления весов.
  • Применяет выбранный алгоритм оптимизации.

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

Проверка качества сети

Для оценки точности сети используется прямой проход без обучения:

var output = net.forward(x);
var predicted_label = output.w.indexOf(Math.max(...output.w));

Ключевые моменты:

  • forward возвращает объект Vol с вероятностями классов.
  • Сравнение предсказанных классов с реальными метками позволяет вычислять точность и контролировать переобучение.

Сохранение и загрузка модели

ConvNetJS поддерживает сериализацию сети в JSON, что удобно для сохранения и последующей загрузки:

var json = net.toJSON();
var net2 = new convnetjs.Net();
net2.fromJSON(json);

Особенности:

  • Полная архитектура и веса сети сохраняются.
  • Можно передавать модели между браузером и сервером.

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

Для повышения качества работы агента применяются:

  • Регуляризация (l2_decay, dropout).
  • Изменение функции активации (relu, tanh, sigmoid).
  • Различные методы оптимизации (SGD с моментом, Adam, Adadelta).
  • Обучение на мини-батчах для стабильности.

Применение этих методов позволяет агенту быстро адаптироваться к сложным входным данным и минимизировать ошибку на тестовых примерах.

Обучение агента на примере среды

Для создания агента, способного принимать решения на основе состояния среды:

  1. Представить состояние среды в виде вектора признаков.
  2. Использовать сеть с соответствующим размером входа.
  3. Прямой проход вычисляет оценку действия или вероятность выбора.
  4. На основе функции вознаграждения обновлять веса через trainer.train.

Пример для Q-обучения через ConvNetJS:

var x = new convnetjs.Vol(state);
var q_values = net.forward(x).w;
var target = q_values.slice();
target[action] = reward + gamma * max_future_q;
trainer.train(x, target);

Ключевые моменты Q-обучения:

  • state преобразуется в Vol.
  • q_values — оценки всех действий.
  • Целевая функция учитывает текущую награду и ожидаемую будущую выгоду.
  • Обновление сети через train позволяет агенту оптимизировать стратегию.

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

ConvNetJS встроенно поддерживает получение значений градиентов и активаций каждого слоя, что позволяет создавать интерактивные графики:

console.log(net.layers[1].filters[0].w);

Такое наблюдение помогает:

  • Отслеживать, как фильтры обучаются на изображениях.
  • Контролировать насыщенность градиентов.
  • Настраивать архитектуру сети для улучшения обучения.