Метод compile: структура и обязательные параметры

Метод compile является ключевым элементом при работе с моделями в Keras.js, обеспечивая подготовку нейронной сети к обучению и последующей оценке. В отличие от Python-версии Keras, Keras.js работает в браузере и использует предобученные модели или модели, загруженные из JSON, но принципы подготовки модели к работе сохраняются.


Основная структура метода compile

В Keras.js метод compile вызывается на объекте модели:

model.compile({
    optimizer: 'adam',
    loss: 'categoricalCrossentropy',
    metrics: ['accuracy']
});

Метод принимает один объект с ключами, определяющими поведение модели во время обучения. Все параметры объекта строго определены, и корректная их настройка критически важна для корректной работы модели.


Обязательные параметры

1. optimizer

Назначение: задаёт алгоритм оптимизации весов модели в процессе обучения.

Типы значений:

  • Строка с названием оптимизатора ('sgd', 'adam', 'rmsprop' и др.).
  • Объект оптимизатора с конкретными параметрами, например:
const adam = new KerasJS.Optimizers.Adam({learningRate: 0.001});
model.compile({optimizer: adam, loss: 'meanSquaredError'});

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

  • optimizer должен соответствовать задачам сети: для классификации чаще используется 'adam', для регрессии — 'sgd' или 'rmsprop'.
  • Неправильный выбор оптимизатора может привести к нестабильному обучению.

2. loss

Назначение: определяет функцию потерь, по которой модель оценивает ошибку предсказаний и обновляет веса.

Типы значений:

  • Строка с названием функции потерь:

    • 'meanSquaredError' — для регрессии.
    • 'categoricalCrossentropy' — для многоклассовой классификации.
    • 'binaryCrossentropy' — для бинарной классификации.
  • Пользовательская функция потерь, реализованная в JS.

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

  • Выбор функции потерь напрямую зависит от задачи.
  • Ошибки при указании функции потерь могут приводить к NaN или некорректным градиентам.

Дополнительные параметры

metrics

Назначение: массив метрик, используемых для оценки производительности модели в процессе обучения и тестирования.

Пример:

model.compile({
    optimizer: 'adam',
    loss: 'categoricalCrossentropy',
    metrics: ['accuracy', 'precision']
});

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

  • Метрики не влияют на обучение напрямую, но помогают отслеживать эффективность модели.
  • В Keras.js поддерживаются базовые метрики (accuracy, precision, recall).

lossWeights

Назначение: задаёт вес для каждой функции потерь при многовыходных моделях.

model.compile({
    optimizer: 'adam',
    loss: ['categoricalCrossentropy', 'meanSquaredError'],
    lossWeights: [0.7, 0.3],
    metrics: ['accuracy']
});

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

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

sampleWeightMode

Назначение: определяет способ взвешивания отдельных образцов данных при обучении. Возможные значения: 'temporal' (для последовательных данных), null (по умолчанию).


Принципы использования compile

  1. Сначала модель, потом compile. Сначала необходимо создать модель и добавить слои, только затем вызывать compile. Попытка компиляции пустой модели вызовет ошибку.

  2. Совместимость optimizer и loss. Нельзя использовать оптимизатор и функцию потерь, предназначенные для несовместимых типов задач. Например, 'binaryCrossentropy' с регрессионным выходом приведёт к некорректным результатам.

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

  4. Предобученные модели. Для Keras.js, который ориентирован на работу с JSON-моделями, метод compile может использоваться и для оценки модели без дальнейшего обучения: достаточно указать optimizer и loss, чтобы получить корректные предсказания.


Примеры корректного использования

Классическая многоклассовая классификация:

const model = new KerasJS.Model({
    filepath: 'model.json',
    gpu: true
});

model.compile({
    optimizer: 'adam',
    loss: 'categoricalCrossentropy',
    metrics: ['accuracy']
});

Регрессия с пользовательским оптимизатором:

const sgd = new KerasJS.Optimizers.SGD({learningRate: 0.01});
model.compile({
    optimizer: sgd,
    loss: 'meanSquaredError',
    metrics: ['mse']
});

Метод compile задаёт фундаментальные правила работы нейронной сети, определяя способ корректировки весов и критерии оценки качества. Понимание обязательных параметров optimizer и loss, а также их взаимосвязи с задачей, является необходимым для стабильного и эффективного обучения моделей в Keras.js.