Метод on_epoch_end класса Callback
Метод on_epoch_end класса Callback вызывается автоматически
в конце каждой эпохи обучения, валидации или предсказания.
Первым параметром метод принимает номер текущей эпохи epoch,
отсчёт начинается с нуля. Вторым параметром передаётся словарь logs,
содержащий значения метрик и потерь на данной эпохе.
Метод не возвращает значений, но может использоваться для логирования,
сохранения весов, изменения скорости обучения или ранней остановки.
Синтаксис
class MyCallback(tf.keras.callbacks.Callback):
def on_epoch_end(self, epoch, logs=None):
pass
Пример
Давайте создадим собственный колбэк, который выводит номер эпохи и значение потерь в конце каждой эпохи:
import tensorflow as tf
tf.random.set_seed(0)
class EpochEndLogger(tf.keras.callbacks.Callback):
def on_epoch_end(self, epoch, logs=None):
logs = logs or {}
print("Epoch", epoch, "- loss:", logs.get("loss"))
model = tf.keras.Sequential([
tf.keras.layers.Dense(1, input_shape=(1,))
])
model.compile(optimizer="sgd", loss="mse")
x = tf.constant([[1.0], [2.0], [3.0], [4.0]])
y = tf.constant([[2.0], [4.0], [6.0], [8.0]])
model.fit(x, y, epochs=2, verbose=0, callbacks=[EpochEndLogger()])
Результат выполнения кода:
Epoch 0 - loss: 20.57639503479004
Epoch 1 - loss: 8.986899375915527
Пример
Давайте создадим колбэк, который сохраняет модель в конце каждой эпохи:
import tensorflow as tf
tf.random.set_seed(0)
class SaveOnEpochEnd(tf.keras.callbacks.Callback):
def on_epoch_end(self, epoch, logs=None):
path = "model_epoch_" + str(epoch) + ".keras"
self.model.save(path)
print("Saved:", path)
model = tf.keras.Sequential([
tf.keras.layers.Dense(1, input_shape=(1,))
])
model.compile(optimizer="sgd", loss="mse")
x = tf.constant([[1.0], [2.0], [3.0]])
y = tf.constant([[2.0], [4.0], [6.0]])
model.fit(x, y, epochs=2, verbose=0, callbacks=[SaveOnEpochEnd()])
Результат выполнения кода:
Saved: model_epoch_0.keras
Saved: model_epoch_1.keras
Пример
Давайте создадим колбэк, который останавливает обучение, если потери стали меньше заданного порога:
import tensorflow as tf
tf.random.set_seed(0)
class StopOnLowLoss(tf.keras.callbacks.Callback):
def on_epoch_end(self, epoch, logs=None):
logs = logs or {}
if logs.get("loss", float("inf")) < 5.0:
print("Stopping at epoch", epoch)
self.model.stop_training = True
model = tf.keras.Sequential([
tf.keras.layers.Dense(1, input_shape=(1,))
])
model.compile(optimizer="sgd", loss="mse")
x = tf.constant([[1.0], [2.0], [3.0]])
y = tf.constant([[2.0], [4.0], [6.0]])
model.fit(x, y, epochs=20, verbose=0, callbacks=[StopOnLowLoss()])
Результат выполнения кода:
"Stopping at epoch 2"
Смотрите также
-
класс
Callback,
который является базовым классом для создания колбэков -
метод
on_epoch_begin,
который вызывается в начале каждой эпохи -
метод
on_train_end,
который вызывается по завершении обучения -
метод
set_model,
который сохраняет ссылку на модель внутри колбэка