Spec-Zone.ru › TensorFlow 2.3

tf.keras.callbacks.ModelCheckpoint

Просмотреть исходный код на GitHub

Обработчик для сохранения модели Keras или весов модели с определённой частотой.

Наследуется от: Callback

Просмотр псевдонимов

Псевдонимы для миграции

См. Руководство по миграции для получения более подробной информации.

tf.compat.v1.keras.callbacks.ModelCheckpoint

tf.keras.callbacks.ModelCheckpoint(
    filepath, monitor='val_loss', verbose=0, save_best_only=False,
    save_weights_only=False, mode='auto', save_freq='epoch', options=None, **kwargs
)

ModelCheckpoint обработчик используется совместно с обучением с помощью model.fit() для сохранения модели или весов (в файле контрольной точки) с определённым интервалом, чтобы модель или веса могли быть загружены позже для продолжения обучения с сохранённого состояния.

Некоторые возможности этого обработчика включают:

  • Сохранять только модель, которая достигла наилучшего результата до сих пор, или сохранять модель в конце каждой эпохи независимо от результата.
  • Определение «лучшего»; какой показатель отслеживать и как его максимизировать или минимизировать.
  • Частоту сохранения. В настоящее время обработчик поддерживает сохранение в конце каждой эпохи или после фиксированного числа обучающих пакетов.
  • Сохранять только веса или всю модель.

Пример:

EPOCHS = 10
checkpoint_filepath = '/tmp/checkpoint'
model_checkpoint_callback = tf.keras.callbacks.ModelCheckpoint(
    filepath=checkpoint_filepath,
    save_weights_only=True,
    monitor='val_acc',
    mode='max',
    save_best_only=True)

# Model weights are saved at the end of every epoch, if it's the best seen
# so far.
model.fit(epochs=EPOCHS, callbacks=[model_checkpoint_callback])

# The model weights (that are considered the best) are loaded into the model.
model.load_weights(checkpoint_filepath)
Аргументы
filepath строка или PathLike, путь для сохранения файла модели. filepath может содержать именованные форматирующие опции, которые будут заполнены значением epoch и ключами в logs (переданные в on_epoch_end). Например, если filepath это weights.{epoch:02d}-{val_loss:.2f}.hdf5, контрольные точки модели будут сохраняться с номером эпохи и значением валидационной ошибки в имени файла.
monitor показатель для отслеживания.
verbose режим подробности, 0 или 1.
save_best_only если save_best_only=True, последняя лучшая модель в соответствии с отслеживаемым показателем не будет перезаписана. Если filepath не содержит форматирующих опций, как {epoch}, тогда filepath будет перезаписана каждой новой лучшей моделью.
mode одно из {auto, min, max}. Если save_best_only=True, решение о перезаписи текущего файла сохранения принимается на основе максимизации или минимизации отслеживаемого показателя. Для val_acc, это должно быть max, для val_loss это должно быть min, и т. д. В режиме auto, направление автоматически выводится из названия отслеживаемого показателя.
save_weights_only если True, то сохраняются только веса модели (model.save_weights(filepath)), иначе сохраняется вся модель (model.save(filepath)).
save_freq 'epoch' или целое число. При использовании 'epoch', обработчик сохраняет модель после каждой эпохи. При использовании целого числа обработчик сохраняет модель в конце этого количества пакетов. Если Model скомпилирована с experimental_steps_per_execution=N, критерии сохранения будут проверяться каждые N-й пакет. Обратите внимание, что если сохранение не выровнено по эпохам, отслеживаемый показатель может быть менее надёжным (он может отражать всего 1 пакет, так как метрики сбрасываются каждую эпоху). По умолчанию 'epoch'.
options Необязательный tf.train.CheckpointOptions объект, если save_weights_only равно true, или необязательный tf.saved_model.SavedOptions объект, если save_weights_only равно false.
**kwargs Дополнительные аргументы для обратной совместимости. Возможный ключ period.

Методы

set_model

Просмотреть исходный код

set_model(
    model
)

set_params

Просмотреть исходный код

set_params(
    params
)

© 2020 The TensorFlow Authors. All rights reserved.
Licensed under the Creative Commons Attribution License 3.0.
Code samples licensed under the Apache 2.0 License.
https://www.tensorflow.org/versions/r2.3/api_docs/python/tf/keras/callbacks/ModelCheckpoint

Spec-Zone.ru

Настройки Оффлайн Что нового Помощь О нас
Spec-Zone .ru
спецификации, руководства, описания, API