Callbacks: onEpochEnd, onBatchEnd, EarlyStopping, ModelCheckpoint

Callbacks в TensorFlow.js представляют собой объекты, позволяющие отслеживать процесс обучения моделей и вмешиваться в него на разных этапах. Они обеспечивают гибкость и контроль над обучением, включая возможность логирования, динамической корректировки параметров, остановки обучения или сохранения модели. Наиболее часто используемыми callback’ами являются onEpochEnd, onBatchEnd, EarlyStopping и ModelCheckpoint.


onEpochEnd и onBatchEnd

onEpochEnd вызывается после завершения каждой эпохи обучения модели. Эпоха — это один полный проход по всем обучающим данным. Основные задачи, решаемые этим callback:

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

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

const history = await model.fit(xs, ys, {
  epochs: 50,
  callbacks: {
    onEpochEnd: (epoch, logs) => {
      console.log(`Эпоха ${epoch + 1}: потеря = ${logs.loss.toFixed(4)}, точность = ${logs.acc?.toFixed(4)}`);
    }
  }
});

onBatchEnd вызывается после каждой мини-партии (batch) данных. Мини-партии используются для более эффективного обучения на больших наборах данных и позволяют обновлять веса модели после каждого небольшого шага. Основные сценарии использования:

  • Тонкое отслеживание процесса обучения на уровне batch.
  • Динамическое логирование или визуализация потерь.
  • Реализация сложных правил изменения параметров обучения.

Пример:

const history = await model.fit(xs, ys, {
  epochs: 5,
  batchSize: 32,
  callbacks: {
    onBatchEnd: async (batch, logs) => {
      console.log(`Batch ${batch + 1}: потеря = ${logs.loss.toFixed(4)}`);
      await tf.nextFrame(); // предотвращение блокировки UI при работе в браузере
    }
  }
});

EarlyStopping

Callback EarlyStopping предназначен для раннего прекращения обучения модели, если метрика качества перестает улучшаться. Это предотвращает переобучение и экономит ресурсы. Основные параметры:

  • monitor – метрика, по которой оценивается улучшение (loss, val_loss, acc и др.).
  • patience – количество эпох без улучшения, после которых обучение останавливается.
  • minDelta – минимальное изменение метрики, чтобы считалось улучшением.
  • mode – режим сравнения метрики ('min' для потерь, 'max' для точности).

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

const earlyStopping = tf.callbacks.earlyStopping({
  monitor: 'val_loss',
  patience: 5,
  minDelta: 0.001,
  mode: 'min'
});

await model.fit(xs, ys, {
  epochs: 100,
  validationSplit: 0.2,
  callbacks: [earlyStopping]
});

В этом примере обучение автоматически завершится, если валидирующая потеря не улучшится на 0.001 в течение 5 эпох.


ModelCheckpoint

ModelCheckpoint позволяет сохранять модель или её веса в процессе обучения. Это полезно для восстановления модели после прерывания обучения или для сохранения лучшей версии модели. Основные параметры:

  • filepath – путь или шаблон имени файла для сохранения.
  • saveWeightsOnly – сохранять только веса (true) или всю модель (false).
  • monitor – метрика, по которой определяется “лучшая” модель.
  • saveBestOnly – сохранять только модель с наилучшей метрикой.

Пример:

const checkpoint = tf.callbacks.modelCheckpoint({
  filepath: 'model-epoch-{epoch:02d}-val_loss-{val_loss:.4f}.json',
  saveWeightsOnly: false,
  monitor: 'val_loss',
  saveBestOnly: true
});

await model.fit(xs, ys, {
  epochs: 50,
  validationSplit: 0.2,
  callbacks: [checkpoint]
});

Этот код сохраняет модель после каждой эпохи, но фактически записывает только версию с наименьшей валидирующей потерей.


Комплексное использование callback’ов

Callback’ы можно комбинировать для максимального контроля над обучением. Например, одновременно логировать процесс, сохранять лучшие модели и останавливать обучение при отсутствии улучшений:

await model.fit(xs, ys, {
  epochs: 100,
  batchSize: 32,
  validationSplit: 0.2,
  callbacks: [
    tf.callbacks.earlyStopping({ monitor: 'val_loss', patience: 10, mode: 'min' }),
    tf.callbacks.modelCheckpoint({
      filepath: 'best-model.json',
      saveBestOnly: true,
      monitor: 'val_loss'
    }),
    {
      onEpochEnd: (epoch, logs) => {
        console.log(`Epoch ${epoch + 1}: loss = ${logs.loss.toFixed(4)}, val_loss = ${logs.val_loss?.toFixed(4)}`);
      }
    }
  ]
});

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


Callbacks в TensorFlow.js позволяют создавать обучающие циклы, сопоставимые с Python-версией, сохраняя при этом преимущества работы непосредственно в браузере или на Node.js, включая интерактивность и контроль за ресурсами.