Spec-Zone.ru › TensorFlow 2.4

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() для сохранения модели или весов (в файле контрольной точки) с определённым интервалом, чтобы модель или веса могли быть загружены позже для продолжения обучения с сохранённого состояния.

Этот обработчик предоставляет несколько вариантов:

  • Сохранять только модель, достигшую наилучшей производительности на данный момент, или сохранять модель в конце каждой эпохи независимо от производительности.
  • Определение «наилучшей» производительности: какой показатель отслеживать и следует ли его максимизировать или минимизировать.
  • Частота сохранения. В настоящее время обработчик поддерживает сохранение в конце каждой эпохи или после определённого количества обучающих партий.
  • Сохранять только веса или всю модель.
Примечание: Если у вас возникнут 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. Примечание:
  • Предшествуйте имени префиксом "val_ для отслеживания метрик валидации.
  • Используйте "loss" или "val_loss" для отслеживания общей потери модели.
  • Если вы указываете метрики как строки, например, "accuracy", передайте ту же строку (с префиксом "val_" или без него).
  • Если вы передаёте объекты metrics.Metric, monitor должно быть установлено в metric.name
  • Если вы не уверены в именах метрик, вы можете проверить содержимое словаря history.history возвращаемого методом history = model.fit()
  • Модели с несколькими выходами добавляют дополнительные префиксы к именам метрик.
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

Spec-Zone.ru

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