Метод on_predict_end класса Callback
Метод on_predict_end принадлежит классу Callback в TensorFlow и вызывается автоматически в конце процесса предсказания, то есть после завершения работы метода predict. Этот метод не принимает обязательных параметров, кроме self, и не возвращает значения. Его можно переопределить в пользовательском классе-наследнике, чтобы выполнить завершающие действия: вывести сообщение, сохранить метрики, закрыть файлы или освободить ресурсы.
Синтаксис
class MyCallback(tf.keras.callbacks.Callback):
def on_predict_end(self, logs=None):
pass
Пример
Давайте создадим простую модель и колбэк, который выводит сообщение в конце предсказания:
import tensorflow as tf
tf.random.set_seed(0)
class PredictEndCallback(tf.keras.callbacks.Callback):
def on_predict_end(self, logs=None):
print("Prediction finished")
model = tf.keras.Sequential([
tf.keras.layers.Dense(1, input_shape=(1,))
])
model.compile(optimizer='sgd', loss='mse')
data = tf.constant([[1.0], [2.0], [3.0]])
callback = PredictEndCallback()
model.predict(data, callbacks=[callback])
Результат выполнения кода:
"Prediction finished"
Пример
Давайте используем параметр logs для сохранения информации о завершении предсказания:
import tensorflow as tf
tf.random.set_seed(0)
class LoggingPredictCallback(tf.keras.callbacks.Callback):
def on_predict_end(self, logs=None):
if logs is None:
logs = {}
logs['predict_status'] = 'completed'
print(logs)
model = tf.keras.Sequential([
tf.keras.layers.Dense(1, input_shape=(1,))
])
model.compile(optimizer='sgd', loss='mse')
data = tf.constant([[1.0], [2.0], [3.0]])
callback = LoggingPredictCallback()
model.predict(data, callbacks=[callback])
Результат выполнения кода:
{'predict_status': 'completed'}
Смотрите также
-
класс
Callback,
который является базовым классом для всех колбэков -
метод
on_predict_begin,
который вызывается в начале предсказания -
метод
on_test_end,
который вызывается в конце тестирования -
метод
on_train_end,
который вызывается в конце обучения