Метод set_model класса Callback
Метод set_model класса Callback привязывает модель Keras к колбэку.
Он вызывается автоматически внутри метода fit перед началом обучения и сохраняет
переданную модель в атрибуте self.model. Благодаря этому колбэк может обращаться
к слоям, весам, оптимизатору и метрикам модели прямо во время тренировки.
Первым параметром метод принимает саму модель, которую нужно привязать.
Самостоятельно вызывать этот метод нужно только при ручном управлении циклом обучения.
Синтаксис
callback.set_model(model)
Пример
Давайте создадим собственный колбэк, переопределим в нём метод set_model
и проверим, что модель успешно привязалась:
import tensorflow as tf
class MyCallback(tf.keras.callbacks.Callback):
def set_model(self, model):
super().set_model(model)
print("model attached:", model.name)
model = tf.keras.Sequential([
tf.keras.layers.Dense(1, input_shape=(1,))
], name="my_model")
callback = MyCallback()
callback.set_model(model)
Результат выполнения кода:
"model attached: my_model"
Пример
Давайте обучим модель с колбэком, который в методе on_epoch_end
обращается к привязанной модели через self.model:
import tensorflow as tf
tf.random.set_seed(0)
class InspectCallback(tf.keras.callbacks.Callback):
def on_epoch_end(self, epoch, logs=None):
weights = self.model.get_weights()
print("epoch:", epoch, "| weights:", weights)
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=[InspectCallback()])
Результат выполнения кода:
epoch: 0 | weights: [array([[0.30214643]], dtype=float32), array([0.01391602], dtype=float32)]
epoch: 1 | weights: [array([[0.67210084]], dtype=float32), array([0.03047276], dtype=float32)]
Смотрите также
-
класс
Callback,
который является базовым классом для всех колбэков -
метод
set_params,
который привязывает параметры обучения к колбэку -
метод
on_train_begin,
который вызывается в начале обучения модели -
метод
on_epoch_end,
который вызывается в конце каждой эпохи обучения