Callbacks в TensorFlow.js представляют собой объекты, позволяющие
отслеживать процесс обучения моделей и вмешиваться в него на разных
этапах. Они обеспечивают гибкость и контроль над обучением, включая
возможность логирования, динамической корректировки параметров,
остановки обучения или сохранения модели. Наиболее часто используемыми
callback’ами являются onEpochEnd, onBatchEnd,
EarlyStopping и ModelCheckpoint.
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) данных. Мини-партии используются для более
эффективного обучения на больших наборах данных и позволяют обновлять
веса модели после каждого небольшого шага. Основные сценарии
использования:
Пример:
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 при работе в браузере
}
}
});
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 позволяет сохранять
модель или её веса в процессе обучения. Это полезно для восстановления
модели после прерывания обучения или для сохранения лучшей версии
модели. Основные параметры:
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’ы можно комбинировать для максимального контроля над обучением. Например, одновременно логировать процесс, сохранять лучшие модели и останавливать обучение при отсутствии улучшений:
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, включая интерактивность и контроль за ресурсами.