РЕПЕТИТОР математика физика информатика
Для школьников и студентов. Подтягивание пробелов. ЦЭ, ЦТ, ОГЭ, ЕГЭ.
Идет набор на ЛЕТО. Жмите для подробностей:)
422 of 824 menu

Класс 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-файл
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить