Метод on_train_end
Метод on_train_end принадлежит классу Callback в TensorFlow.
Он автоматически вызывается фреймворком в момент завершения процесса обучения модели,
то есть после того, как отработаны все эпохи. Этот метод удобно использовать для
финальной очистки ресурсов, сохранения итоговых метрик или вывода сводной информации.
Метод не принимает обязательных параметров, кроме self, и не возвращает значений,
однако вы можете переопределить его в своём классе-наследнике.
Синтаксис
class MyCallback(tf.keras.callbacks.Callback):
def on_train_end(self, logs=None):
# логика по завершении обучения
Пример
Давайте создадим простой колбэк, который выводит сообщение об окончании обучения и количество прошедших эпох:
import tensorflow as tf
tf.random.set_seed(0)
class TrainEndLogger(tf.keras.callbacks.Callback):
def on_train_end(self, logs=None):
print("Training finished")
print(f"Epochs completed: {self.params['epochs']}")
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, verbose=0, callbacks=[TrainEndLogger()])
Результат выполнения кода:
"Training finished"
"Epochs completed: 3"
Пример
Метод можно использовать для сохранения модели после завершения обучения.
Создадим колбэк, который сохраняет модель в файл model.keras:
import tensorflow as tf
tf.random.set_seed(0)
class SaveOnTrainEnd(tf.keras.callbacks.Callback):
def on_train_end(self, logs=None):
self.model.save('model.keras')
print("Model saved to model.keras")
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=[SaveOnTrainEnd()])
Результат выполнения кода:
"Model saved to model.keras"
Смотрите также
-
класс
Callback,
который является базовым для создания колбэков -
метод
on_train_begin,
который вызывается перед началом обучения -
метод
on_epoch_end,
который вызывается в конце каждой эпохи -
метод
set_model,
который устанавливает модель для колбэка