Класс ModelCheckpoint
Класс ModelCheckpoint представляет собой колбэк,
который сохраняет модель или ее веса после каждой
эпохи обучения. Первым параметром передается путь
к файлу для сохранения. Вторым параметром можно
указать метрику для отслеживания. Третьим - режим
сравнения. Также можно сохранять только лучшие веса
через параметр save_best_only.
Синтаксис
tf.keras.callbacks.ModelCheckpoint(
filepath,
monitor="val_loss",
verbose=0,
save_best_only=False,
save_weights_only=False,
mode="auto",
save_freq="epoch"
)
Пример
Давайте создадим простую модель и сохраним ее
во время обучения с помощью ModelCheckpoint:
import tensorflow as tf
tf.random.set_seed(0)
x = tf.constant([[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]])
y = tf.constant([[1.0], [0.0], [1.0]])
model = tf.keras.Sequential([
tf.keras.layers.Dense(4, activation="relu"),
tf.keras.layers.Dense(1, activation="sigmoid")
])
model.compile(optimizer="adam", loss="binary_crossentropy")
checkpoint = tf.keras.callbacks.ModelCheckpoint(
filepath="model.keras",
monitor="loss",
save_best_only=True,
verbose=1
)
model.fit(x, y, epochs=3, callbacks=[checkpoint])
Результат выполнения кода:
Epoch 1/3
1/1 ━━━━━━━━━━━━━━━━━━━━ 0s 200ms/step - loss: 0.6931
Epoch 1: loss improved from inf to 0.69315, saving model to model.keras
Epoch 2/3
1/1 ━━━━━━━━━━━━━━━━━━━━ 0s 30ms/step - loss: 0.6890
Epoch 2: loss improved from 0.69315 to 0.68901, saving model to model.keras
Epoch 3/3
1/1 ━━━━━━━━━━━━━━━━━━━━ 0s 28ms/step - loss: 0.6848
Epoch 3: loss improved from 0.68901 to 0.68478, saving model to model.keras
Пример
Давайте сохраним только веса модели с
помощью параметра save_weights_only:
import tensorflow as tf
tf.random.set_seed(0)
x = tf.constant([[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]])
y = tf.constant([[1.0], [0.0], [1.0]])
model = tf.keras.Sequential([
tf.keras.layers.Dense(4, activation="relu"),
tf.keras.layers.Dense(1, activation="sigmoid")
])
model.compile(optimizer="adam", loss="binary_crossentropy")
checkpoint = tf.keras.callbacks.ModelCheckpoint(
filepath="weights.weights.h5",
save_weights_only=True,
verbose=1
)
model.fit(x, y, epochs=2, callbacks=[checkpoint])
model.load_weights("weights.weights.h5")
res = model.predict(x)
print(res)
Результат выполнения кода:
Epoch 1/2
1/1 ━━━━━━━━━━━━━━━━━━━━ 0s 200ms/step - loss: 0.6931
Epoch 1: saving model to weights.weights.h5
Epoch 2/2
1/1 ━━━━━━━━━━━━━━━━━━━━ 0s 30ms/step - loss: 0.6890
Epoch 2: saving model to weights.weights.h5
1/1 ━━━━━━━━━━━━━━━━━━━━ 0s 50ms/step
[[0.5000001]
[0.4999999]
[0.5000001]]
Смотрите также
-
класс
EarlyStopping,
который останавливает обучение при отсутствии улучшений -
класс
ReduceLROnPlateau,
который уменьшает скорость обучения при остановке метрики -
класс
TensorBoard,
который записывает логи для визуализации обучения -
класс
CSVLogger,
который сохраняет историю обучения в CSV-файл