Метод on_train_begin класса Callback
Метод on_train_begin класса Callback вызывается один раз в самом начале обучения модели, до обработки первой эпохи. Метод принимает один параметр logs - словарь, содержащий информацию о текущем состоянии обучения. Этот метод удобно использовать для инициализации переменных, сброса счетчиков, логирования времени старта или подготовки внешних ресурсов перед тренировкой.
Синтаксис
class MyCallback(tf.keras.callbacks.Callback):
def on_train_begin(self, logs=None):
pass
Пример
Давайте создадим простой колбэк, который выводит сообщение в начале обучения:
import tensorflow as tf
tf.random.set_seed(0)
class StartCallback(tf.keras.callbacks.Callback):
def on_train_begin(self, logs=None):
print("Training started")
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]])
y = tf.constant([[2.0], [4.0], [6.0], [8.0], [10.0]])
model.fit(x, y, epochs=2, callbacks=[StartCallback()], verbose=0)
Результат выполнения кода:
Training started
Пример
Давайте используем параметр logs для вывода доступной информации в начале обучения:
import tensorflow as tf
tf.random.set_seed(0)
class LogsCallback(tf.keras.callbacks.Callback):
def on_train_begin(self, logs=None):
print("logs:", logs)
print("params:", self.params)
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]])
y = tf.constant([[2.0], [4.0], [6.0], [8.0], [10.0]])
model.fit(x, y, epochs=2, callbacks=[LogsCallback()], verbose=0)
Результат выполнения кода:
logs: {}
params: {'verbose': 0, 'epochs': 2, 'steps': 5}
Пример
Давайте применим on_train_begin для сброса счетчика эпох перед обучением:
import tensorflow as tf
tf.random.set_seed(0)
class CounterCallback(tf.keras.callbacks.Callback):
def __init__(self):
super().__init__()
self.epoch_count = 0
def on_train_begin(self, logs=None):
self.epoch_count = 0
print("Counter reset to", self.epoch_count)
def on_epoch_end(self, epoch, logs=None):
self.epoch_count += 1
print("Epoch count:", self.epoch_count)
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]])
y = tf.constant([[2.0], [4.0], [6.0], [8.0], [10.0]])
model.fit(x, y, epochs=3, callbacks=[CounterCallback()], verbose=0)
Результат выполнения кода:
Counter reset to 0
Epoch count: 1
Epoch count: 2
Epoch count: 3
Смотрите также
-
класс
Callback,
который является базовым классом для создания колбэков -
метод
on_train_end,
который вызывается в конце обучения модели -
метод
on_epoch_begin,
который вызывается в начале каждой эпохи -
метод
set_params,
который устанавливает параметры обучения для колбэка