Метод on_batch_end класса Callback
Метод on_batch_end класса Callback
вызывается автоматически в конце каждого батча
во время обучения, оценки или предсказания модели.
Метод принимает один параметр batch -
номер текущего батча (целое число, начиная с 0),
а также может принимать параметр logs -
словарь с метриками текущего батча.
Метод не возвращает значения и используется
для реализации пользовательской логики:
логирования, изменения скорости обучения,
сохранения контрольных точек и других действий.
Синтаксис
class MyCallback(tf.keras.callbacks.Callback):
def on_batch_end(self, batch, logs=None):
# logic at the end of each batch
Пример
Давайте создадим простой колбэк, который выводит номер завершенного батча и значение функции потерь:
import tensorflow as tf
import numpy as np
tf.random.set_seed(0)
class BatchLogger(tf.keras.callbacks.Callback):
def on_batch_end(self, batch, logs=None):
logs = logs or {}
print(f"batch {batch} ended, loss: {logs.get('loss')}")
x = np.random.rand(100, 4)
y = np.random.rand(100, 1)
model = tf.keras.Sequential([
tf.keras.layers.Dense(1, input_shape=(4,))
])
model.compile(optimizer='sgd', loss='mse')
model.fit(x, y, epochs=1, batch_size=32, callbacks=[BatchLogger()], verbose=0)
Результат выполнения кода:
"batch 0 ended, loss: ..."
"batch 1 ended, loss: ..."
"batch 2 ended, loss: ..."
"batch 3 ended, loss: ..."
Пример
Давайте сделаем колбэк, который снижает скорость обучения после третьего батча:
import tensorflow as tf
import numpy as np
tf.random.set_seed(0)
class LrReducer(tf.keras.callbacks.Callback):
def on_batch_end(self, batch, logs=None):
if batch == 2:
lr = float(tf.keras.backend.get_value(self.model.optimizer.learning_rate))
tf.keras.backend.set_value(self.model.optimizer.learning_rate, lr * 0.5)
print(f"learning rate reduced to {lr * 0.5}")
x = np.random.rand(100, 4)
y = np.random.rand(100, 1)
model = tf.keras.Sequential([
tf.keras.layers.Dense(1, input_shape=(4,))
])
model.compile(optimizer=tf.keras.optimizers.SGD(learning_rate=0.1), loss='mse')
model.fit(x, y, epochs=1, batch_size=32, callbacks=[LrReducer()], verbose=0)
Результат выполнения кода:
"learning rate reduced to 0.05"
Смотрите также
-
класс
Callback,
который является базовым классом для всех колбэков -
метод
on_batch_begin,
который вызывается в начале каждого батча -
метод
on_epoch_end,
который вызывается в конце каждой эпохи -
метод
set_model,
который устанавливает модель для колбэка