tf.keras.callbacks.ModelCheckpoint
| Просмотреть исходный код на GitHub |
Обработчик для сохранения модели Keras или весов модели с определённой частотой.
Наследуется от: Callback
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