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() для сохранения модели или весов (в файле контрольной точки) с определённым интервалом, чтобы модель или веса могли быть загружены позже для продолжения обучения с сохранённого состояния.
Этот обработчик предоставляет несколько вариантов:
- Сохранять только модель, достигшую наилучшей производительности на данный момент, или сохранять модель в конце каждой эпохи независимо от производительности.
- Определение «наилучшей» производительности: какой показатель отслеживать и следует ли его максимизировать или минимизировать.
- Частота сохранения. В настоящее время обработчик поддерживает сохранение в конце каждой эпохи или после определённого количества обучающих партий.
- Сохранять только веса или всю модель.
Примечание: Если у вас возникнутWARNING:tensorflow:Can save best model only with <name> available, skipping, см. описание аргументаmonitorдля получения подробной информации о том, как это исправить.
Пример:
model.compile(loss=..., optimizer=...,
metrics=['accuracy'])
EPOCHS = 10
checkpoint_filepath = '/tmp/checkpoint'
model_checkpoint_callback = tf.keras.callbacks.ModelCheckpoint(
filepath=checkpoint_filepath,
save_weights_only=True,
monitor='val_accuracy',
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 | Имя метрики для отслеживания. Обычно метрики задаются методом Model.compile. Примечание:
|
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', обработчик сохраняет модель после каждой эпохи. При использовании целого числа обработчик сохраняет модель в конце указанного количества партий. Если модель скомпилирована с steps_per_execution=N, критерии сохранения будут проверяться каждые N-ую партию. Обратите внимание, что если сохранение не выровнено по эпохам, отслеживаемая метрика может быть потенциально менее надёжной (она может отражать всего 1 партию, так как метрики сбрасываются в конце каждой эпохи). По умолчанию 'epoch'. |
options | необъект tf.train.CheckpointOptions, если save_weights_only равно True, или необъект tf.saved_model.SaveOptions, если 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.4/api_docs/python/tf/keras/callbacks/ModelCheckpoint