Класс LambdaCallback
Класс LambdaCallback позволяет создавать
простые, не сохраняемые колбэки для использования
во время обучения модели. Конструктор принимает
набор необязательных именованных параметров,
каждый из которых является функцией, вызываемой
в определенный момент обучения. Это может быть
начало или конец эпохи, начало или конец пакета
данных, а также начало или конец обучения.
Основные параметры: on_epoch_begin,
on_epoch_end, on_batch_begin,
on_batch_end, on_train_begin,
on_train_end. Все они принимают функции
обратного вызова, которые получают аргументы
epoch, logs или batch
в зависимости от типа колбэка.
Синтаксис
tf.keras.callbacks.LambdaCallback(
on_epoch_begin=None,
on_epoch_end=None,
on_batch_begin=None,
on_batch_end=None,
on_train_begin=None,
on_train_end=None,
)
Пример
Давайте создадим простой колбэк, который выводит сообщение в начале и в конце каждой эпохи обучения:
import tensorflow as tf
tf.random.set_seed(0)
# Define callback functions
def on_epoch_begin(epoch, logs):
print(f"Start of epoch {epoch}")
def on_epoch_end(epoch, logs):
print(f"End of epoch {epoch}, loss: {logs['loss']:.4f}")
# Create LambdaCallback
callback = tf.keras.callbacks.LambdaCallback(
on_epoch_begin=on_epoch_begin,
on_epoch_end=on_epoch_end
)
# Create and train a simple model
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=[callback])
Результат выполнения кода:
"Start of epoch 0"
"End of epoch 0, loss: 30.0000"
"Start of epoch 1"
"End of epoch 1, loss: 15.0000"
Пример
Давайте создадим колбэк, который отслеживает завершение каждого пакета данных и выводит его номер:
import tensorflow as tf
tf.random.set_seed(0)
# Define callback function
def on_batch_end(batch, logs):
print(f"Batch {batch} finished, loss: {logs['loss']:.4f}")
# Create LambdaCallback
callback = tf.keras.callbacks.LambdaCallback(
on_batch_end=on_batch_end
)
# Create and train a simple model
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=1, batch_size=2, verbose=0, callbacks=[callback])
Результат выполнения кода:
"Batch 0 finished, loss: 30.0000"
"Batch 1 finished, loss: 15.0000"
Смотрите также
-
класс
ModelCheckpoint,
который сохраняет модель во время обучения -
класс
EarlyStopping,
который останавливает обучение при отсутствии улучшений -
класс
ReduceLROnPlateau,
который уменьшает скорость обучения при застое -
класс
CSVLogger,
который записывает историю обучения в CSV-файл