tf.keras.callbacks.experimental.BackupAndRestore
Обработчик для резервного копирования и восстановления состояния обучения.
Унаследован от: Callback
tf.keras.callbacks.experimental.BackupAndRestore(
backup_dir
)
Обработчик callback предназначен для восстановления после прерываний, произошедших во время выполнения 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.4/api_docs/python/tf/keras/callbacks/experimental/BackupAndRestore