Метод on_predict_begin
Метод on_predict_begin класса Callback вызывается один раз
в начале процесса предсказания модели при вызове метода
predict. Этот метод не принимает обязательных параметров,
кроме self, и не возвращает значений. Его можно переопределить
в собственном классе-наследнике Callback, чтобы выполнить
подготовительные действия перед предсказанием: например, инициализировать
счетчики, зафиксировать время старта или выделить ресурсы.
Синтаксис
class MyCallback(tf.keras.callbacks.Callback):
def on_predict_begin(self, logs=None):
pass
Пример
Давайте создадим простой колбэк, который выводит сообщение
в начале предсказания, и применим его при вызове predict:
import tensorflow as tf
class PredictBeginCallback(tf.keras.callbacks.Callback):
def on_predict_begin(self, logs=None):
print("predict begin")
model = tf.keras.Sequential([
tf.keras.layers.Dense(1, input_shape=(1,))
])
model.compile(optimizer="sgd", loss="mse")
callback = PredictBeginCallback()
res = model.predict([1.0, 2.0, 3.0], callbacks=[callback], verbose=0)
print(res.shape)
Результат выполнения кода:
"predict begin"
(3, 1)
Пример
Давайте используем on_predict_begin для сохранения времени
старта предсказания и последующего вывода длительности
в on_predict_end:
import tensorflow as tf
import time
class TimePredictCallback(tf.keras.callbacks.Callback):
def on_predict_begin(self, logs=None):
self.start_time = time.time()
print("start time saved")
def on_predict_end(self, logs=None):
elapsed = time.time() - self.start_time
print("elapsed:", round(elapsed, 4))
model = tf.keras.Sequential([
tf.keras.layers.Dense(1, input_shape=(1,))
])
model.compile(optimizer="sgd", loss="mse")
callback = TimePredictCallback()
res = model.predict([1.0, 2.0, 3.0], callbacks=[callback], verbose=0)
print(res.shape)
Результат выполнения кода:
"start time saved"
"elapsed: 0.0"
(3, 1)
Смотрите также
-
класс
Callback,
который является базовым классом для всех колбэков -
метод
on_predict_end,
который вызывается в конце предсказания -
метод
on_train_begin,
который вызывается в начале обучения -
метод
on_test_begin,
который вызывается в начале тестирования