tf.keras.callbacks.experimental.BackupAndRestore
Обратный вызов для резервного копирования и восстановления состояния обучения.
Наследуется от: Callback
tf.keras.callbacks.experimental.BackupAndRestore(
backup_dir
)
BackupAndRestore Обратный вызов предназначен для восстановления после прерываний, произошедших в середине выполнения model.fit, создавая резервные копии состояния обучения в временном файле контрольной точки (на основе TF CheckpointManager) в конце каждой эпохи. Если обучение возобновлено до завершения, состояние обучения и модель восстанавливаются до последнего сохранённого состояния в начале нового выполнения model.fit(). Ответственность за перезапуск задач лежит на пользователе. Этот обратный вызов важен для механизма резервного копирования и восстановления с целью повышения отказоустойчивости. При этом ожидается, что модель, восстанавливаемая из предыдущей контрольной точки, будет такой же, как и та, которая использовалась для резервного копирования. Если пользователь изменяет аргументы, передаваемые в compile или fit, сохранённая для повышения отказоустойчивости контрольная точка может стать недействительной.
Примечание:
- Этот обратный вызов несовместим с отключением выполнения в режиме eager.
- Контрольная точка сохраняется в конце каждой эпохи. При восстановлении мы пересчитаем любую частичную работу из незавершенной эпохи, в которой обучение было возобновлено (поэтому проделанная работа до прерывания не повлияет на конечное состояние модели).
- Это работает как для однопроцессорного, так и для многопроцессорного режима. В настоящий момент поддерживаются только 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.3/api_docs/python/tf/keras/callbacks/experimental/BackupAndRestore