Класс Callback
Класс Callback является базовым классом для всех колбэков в TensorFlow.
Колбэк - это объект, который может выполнять определенные действия
на различных этапах жизненного цикла модели: в начале и конце обучения,
в начале и конце каждой эпохи, в начале и конце каждого батча,
а также при тестировании и предсказании. Для создания собственного
колбэка необходимо унаследоваться от класса Callback
и переопределить нужные методы. Класс не принимает обязательных
параметров при инициализации, однако пользовательские колбэки
могут принимать собственные параметры через конструктор.
Синтаксис
tf.keras.callbacks.Callback()
Пример
Давайте создадим простой колбэк, который выводит сообщение в начале и в конце каждой эпохи обучения:
import tensorflow as tf
class MyCallback(tf.keras.callbacks.Callback):
def on_epoch_begin(self, epoch, logs=None):
print("Epoch started:", epoch)
def on_epoch_end(self, epoch, logs=None):
print("Epoch ended:", epoch)
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, callbacks=[MyCallback()], verbose=0)
Результат выполнения кода:
"Epoch started: 0"
"Epoch ended: 0"
"Epoch started: 1"
"Epoch ended: 1"
Пример
Давайте создадим колбэк, который сохраняет номер эпохи в собственный атрибут и выводит его после завершения обучения:
import tensorflow as tf
class EpochLogger(tf.keras.callbacks.Callback):
def __init__(self):
super().__init__()
self.epochs_seen = 0
def on_epoch_end(self, epoch, logs=None):
self.epochs_seen += 1
logger = EpochLogger()
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=3, callbacks=[logger], verbose=0)
print("Epochs seen:", logger.epochs_seen)
Результат выполнения кода:
"Epochs seen: 3"
Смотрите также
-
метод
on_epoch_begin,
который вызывается в начале каждой эпохи -
метод
on_epoch_end,
который вызывается в конце каждой эпохи -
метод
on_train_begin,
который вызывается в начале обучения -
метод
on_train_end,
который вызывается в конце обучения