Spec-Zone.ru › TensorFlow 2.4

tf.keras.callbacks.experimental.BackupAndRestore

Обработчик для резервного копирования и восстановления состояния обучения.

Унаследован от: Callback

tf.keras.callbacks.experimental.BackupAndRestore(
    backup_dir
)

Обработчик callback предназначен для восстановления после прерываний, произошедших во время выполнения model.fit, создавая резервные копии состояния обучения в временном файле контрольной точки (на основе TF CheckpointManager) в конце каждой эпохи. Если обучение было прервано до завершения, состояние обучения и модель восстанавливаются до последнего сохранённого состояния в начале нового выполнения model.fit(). Пользователь отвечает за перезапуск задач. Этот обработчик важен для механизма резервного копирования и восстановления для обеспечения отказоустойчивости. Модель, которая должна быть восстановлена из предыдущей контрольной точки, должна быть такой же, как та, которая использовалась для создания резервной копии. Если пользователь изменяет аргументы, передаваемые в compile или fit, сохранённая контрольная точка для отказоустойчивости может стать недействительной.

Примечание:

  1. Этот обработчик несовместим с отключением выполнения eager.
  2. Контрольная точка сохраняется в конце каждой эпохи. При восстановлении мы повторим любые частичные работы из незавершенной эпохи, в которой обучение было перезапущено (поэтому работа, выполненная до прерывания, не повлияет на конечное состояние модели).
  3. Это работает как для однопоточного, так и для многопоточного режима; в настоящее время поддерживаются только MirroredStrategy и MultiWorkerMirroredStrategy.

Пример:

class InterruptingCallback(tf.keras.callbacks.Callback):
  def on_epoch_begin(self, epoch, logs=None):
    if epoch == 4:
      raise RuntimeError('Interrupting!')
callback = tf.keras.callbacks.experimental.BackupAndRestore(
backup_dir="/tmp")
model = tf.keras.models.Sequential([tf.keras.layers.Dense(10)])
model.compile(tf.keras.optimizers.SGD(), loss='mse')
try:
  model.fit(np.arange(100).reshape(5, 20), np.zeros(5), epochs=10,
            batch_size=1, callbacks=[callback, InterruptingCallback()],
            verbose=0)
except:
  pass
history = model.fit(np.arange(100).reshape(5, 20), np.zeros(5), epochs=10,
            batch_size=1, callbacks=[callback], verbose=0)
# Only 6 more epochs are run, since first trainning got interrupted at
# zero-indexed epoch 4, second training will continue from 4 to 9.
len(history.history['loss'])
6
Аргументы
backup_dir Строка, путь для сохранения файла модели. Это каталог, в котором система хранит временные файлы для восстановления модели после неожиданного завершения задач. Каталог нельзя повторно использовать для хранения других контрольных точек, например, обработчиком BackupAndRestore другого обучения или другим обработчиком (ModelCheckpoint) того же обучения.

Методы

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/experimental/BackupAndRestore

Spec-Zone.ru

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