Spec-Zone.ru › TensorFlow 2.9

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,
    initial_value_threshold=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 = os.path.join(working_dir, 'ckpt', file_name). filepath может содержать именованные параметры форматирования, которые будут заполнены значениями epoch и ключами в logs (переданными в on_epoch_end). Например: если filepath равно weights.{epoch:02d}-{val_loss:.2f}.hdf5, тогда контрольные точки модели будут сохранены с номером эпохи и значением валидационного потерь в имени файла. Директория в пути filepath не должна повторно использоваться другими обработчиками, чтобы избежать конфликтов.
monitor Название метрики для отслеживания. Обычно метрики устанавливаются методом Model.compile. Обратите внимание:
  • Добавьте префикс "val_" к имени, чтобы отслеживать валидационные метрики.
  • Используйте "loss" или "total_loss", чтобы отслеживать общие потери модели.
  • Если вы указываете метрики в виде строк, как "accuracy", передавайте ту же строку (с префиксом "val_" или без него).

  • Если вы передаете объекты метрики metrics.Metric, monitor должен быть установлен на metric.name

  • Если вы не уверены в именах метрик, вы можете проверить содержимое словаря history.history возвращаемого history = model.fit()

  • Модели с несколькими выходами добавляют дополнительные префиксы к именам метрик.

verbose Режим отображения, 0 или 1. Режим 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, режим устанавливается в max если отслеживаемые значения метрики - 'acc' или начинаются с 'fmeasure', и в min для остальных значений.
save_weights_only если True, тогда сохраняются только веса модели (model.save_weights(filepath)), иначе сохраняется вся модель (model.save(filepath)).
save_freq 'epoch' или целое число. При использовании 'epoch', обработчик сохраняет модель после каждой эпохи. При использовании целого числа, обработчик сохраняет модель в конце этого количества батчей. Если Model скомпилирован с steps_per_execution=N, то критерии сохранения будут проверяться каждые N-е батчи. Обратите внимание, что если сохранение не привязано к эпохам, отслеживаемая метрика может быть потенциально менее надёжной (может отражать всего 1 батч, так как метрики сбрасываются в каждой эпохе). По умолчанию 'epoch'.
options Необязательный объект tf.train.CheckpointOptions, если save_weights_only равно true, или необязательный объект tf.saved_model.SaveOptions, если save_weights_only равно false.
initial_value_threshold Начальное значение метрики для отслеживания с плавающей точкой. Применимо только если save_best_value=True. Перезаписывает только сохранённые весовые коэффициенты модели, если производительность текущей модели лучше, чем это значение.
**kwargs Дополнительные аргументы для обратной совместимости. Возможный ключ period.

Методы

set_model

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

set_model(
    model
)

set_params

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

set_params(
    params
)

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

Spec-Zone.ru

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