Метод on_batch_begin класса Callback
Метод on_batch_begin принадлежит классу Callback и вызывается в начале обработки каждого батча.
Первым параметром метод принимает номер текущего батча batch, а вторым - словарь logs с дополнительной информацией.
Метод не возвращает значений, но может использоваться для логирования, изменения состояния или ранней остановки.
Синтаксис
class MyCallback(tf.keras.callbacks.Callback):
def on_batch_begin(self, batch, logs=None):
pass
Пример
Давайте создадим колбэк, который выводит номер батча в начале его обработки:
import tensorflow as tf
tf.random.set_seed(0)
class BatchBeginCallback(tf.keras.callbacks.Callback):
def on_batch_begin(self, batch, logs=None):
print("Begin batch:", batch)
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, callbacks=[BatchBeginCallback()], verbose=0)
Результат выполнения кода:
"Begin batch: 0"
"Begin batch: 1"
Пример
Давайте создадим колбэк, который подсчитывает количество батчей за эпоху:
import tensorflow as tf
tf.random.set_seed(0)
class BatchCounter(tf.keras.callbacks.Callback):
def __init__(self):
super().__init__()
self.batch_count = 0
def on_batch_begin(self, batch, logs=None):
self.batch_count += 1
def on_epoch_end(self, epoch, logs=None):
print("Total batches in epoch:", self.batch_count)
self.batch_count = 0
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], [5.0], [6.0]])
y = tf.constant([[2.0], [4.0], [6.0], [8.0], [10.0], [12.0]])
model.fit(x, y, epochs=2, batch_size=2, callbacks=[BatchCounter()], verbose=0)
Результат выполнения кода:
"Total batches in epoch: 3"
"Total batches in epoch: 3"
Смотрите также
-
класс
Callback,
который является базовым классом для всех колбэков -
метод
on_batch_end,
который вызывается в конце обработки каждого батча -
метод
on_epoch_begin,
который вызывается в начале каждой эпохи -
метод
on_train_begin,
который вызывается в начале обучения модели