Проблема взрывного градиента

Brain.js — это JavaScript-библиотека для создания и обучения нейронных сетей. Она предоставляет удобный интерфейс для работы с различными типами сетей: от простых полносвязных до рекуррентных и LSTM. Основная цель библиотеки — облегчить использование нейронных сетей в веб- и серверных приложениях.

Нейронная сеть в Brain.js строится как последовательность слоев нейронов, соединённых между собой весами. Каждое соединение имеет вес, который изменяется в процессе обучения с помощью алгоритма обратного распространения ошибки.

const brain = require('brain.js');
const net = new brain.NeuralNetwork();

Типы нейронных сетей

  1. Feedforward Neural Network (FFNN) Простая многослойная сеть прямого распространения. Используется для классификации и регрессии. В Brain.js её реализует NeuralNetwork.

  2. Recurrent Neural Network (RNN) Сеть с обратными связями, позволяющая учитывать предыдущие состояния. Реализуется как RecurrentNeuralNetwork.

  3. Long Short-Term Memory (LSTM) Модификация RNN, способная сохранять долгосрочные зависимости в последовательностях. В Brain.js доступна как LSTM и LSTMTimeStep.

Структура данных для обучения

Данные для обучения подаются в виде массива объектов:

const trainingData = [
  { input: [0, 0], output: [0] },
  { input: [0, 1], output: [1] },
  { input: [1, 0], output: [1] },
  { input: [1, 1], output: [0] }
];

input и output могут быть массивами чисел или объектами с ключами и значениями. Для рекуррентных сетей используются последовательности.

Настройка сети и параметры обучения

Нейронная сеть создаётся с возможностью задания конфигурации:

const net = new brain.NeuralNetwork({
  hiddenLayers: [3, 3],
  activation: 'sigmoid', 
  learningRate: 0.01
});

Ключевые параметры:

  • hiddenLayers — массив, задающий количество нейронов в скрытых слоях.
  • activation — функция активации (sigmoid, relu, tanh).
  • learningRate — скорость обучения.
  • iterations — максимальное количество итераций обучения.
  • errorThresh — порог ошибки, после которого обучение останавливается.

Проблема взрывного градиента

Взрывной градиент возникает при обучении глубоких нейронных сетей или рекуррентных сетей, когда значения градиентов становятся слишком большими. Это приводит к нестабильной тренировке: веса начинают изменяться катастрофически быстро, и сеть перестаёт сходиться.

Причины возникновения

  1. Большие веса и глубокие слои В многослойных сетях производная функции активации на некоторых слоях может быть больше единицы, и при обратном распространении ошибка умножается на эти производные на каждом слое, что вызывает экспоненциальный рост градиента.

  2. Функции активации с резкими кривыми Sigmoid и tanh в экстремальных областях имеют очень маленькие или очень большие производные, что может усиливать эффект взрывного градиента.

  3. Длинные последовательности в RNN В рекуррентных сетях ошибка распространяется через большое количество временных шагов, и градиенты на ранних шагах могут стать чрезвычайно большими.

Методы борьбы

  1. Ограничение градиента (Gradient Clipping) В Brain.js можно вручную нормализовать градиенты после обратного распространения, чтобы их величина не превышала заданный порог.

  2. Использование функций активации с мягкими градиентами ReLU и Leaky ReLU уменьшают вероятность взрыва градиента по сравнению с sigmoid и tanh.

  3. Инициализация весов Использование случайной инициализации с малыми значениями помогает ограничить рост градиентов на начальных этапах обучения.

  4. Регуляризация Применение L1/L2-регуляризации или Dropout уменьшает экстремальные изменения весов.

  5. Разбиение последовательностей Для LSTM или RNN длинные последовательности можно делить на более короткие участки, чтобы уменьшить количество шагов, через которые распространяется градиент.

Пример корректного обучения LSTM с предотвращением взрывного градиента

const net = new brain.recurrent.LSTMTimeStep({
  hiddenLayers: [10, 10],
  learningRate: 0.005
});

const trainingData = [
  [1, 2, 3, 4],
  [2, 3, 4, 5],
  [3, 4, 5, 6]
];

net.train(trainingData, {
  iterations: 2000,
  errorThresh: 0.01,
  log: true,
  logPeriod: 100
});

// Ограничение значений предсказаний
const output = net.run([4, 5, 6]);
const clippedOutput = Math.min(Math.max(output, -10), 10);

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

Отладка и визуализация

Brain.js поддерживает экспорт весов и конфигурации сети для анализа. Визуализация распределения весов и активаций на каждом слое помогает выявить признаки взрывного градиента и скорректировать структуру сети до того, как начнётся обучение на больших данных.

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