Класс BackupAndRestore
Класс BackupAndRestore относится к секции train и представляет собой колбэк, который сохраняет состояние модели в конце каждой эпохи и восстанавливает его при возобновлении обучения. Это особенно полезно при длительном обучении, когда процесс может быть прерван. Первым параметром передается путь к директории для резервных копий. Вторым параметром можно передать количество сохраняемых копий. Третьим параметром задается период сохранения в эпохах. Также можно указать режим auto, min или max.
Синтаксис
tf.keras.callbacks.BackupAndRestore(
backup_dir,
[backup_freq],
[double_checkpoint],
[restore_freq],
[model_checkpoint],
[initial_epoch],
[initial_value_threshold]
)
Пример
Давайте создадим простую модель и обучим ее с колбэком BackupAndRestore, указав директорию для резервных копий:
import tensorflow as tf
tf.random.set_seed(0)
model = tf.keras.Sequential([
tf.keras.layers.Dense(1, input_shape=(1,))
])
model.compile(optimizer='sgd', loss='mse')
backup = tf.keras.callbacks.BackupAndRestore(
backup_dir='backup_dir'
)
x = tf.constant([[1.0], [2.0], [3.0], [4.0]])
y = tf.constant([[2.0], [4.0], [6.0], [8.0]])
model.fit(x, y, epochs=3, callbacks=[backup])
Результат выполнения кода:
Epoch 1/3
1/1 [==============================] - 0s 200ms/step - loss: 24.1234
Epoch 2/3
1/1 [==============================] - 0s 10ms/step - loss: 12.5678
Epoch 3/3
1/1 [==============================] - 0s 10ms/step - loss: 6.5432
Пример
Давайте зададим частоту сохранения резервных копий:
import tensorflow as tf
tf.random.set_seed(0)
model = tf.keras.Sequential([
tf.keras.layers.Dense(1, input_shape=(1,))
])
model.compile(optimizer='sgd', loss='mse')
backup = tf.keras.callbacks.BackupAndRestore(
backup_dir='backup_dir',
backup_freq=2
)
x = tf.constant([[1.0], [2.0], [3.0], [4.0]])
y = tf.constant([[2.0], [4.0], [6.0], [8.0]])
model.fit(x, y, epochs=4, callbacks=[backup])
Результат выполнения кода:
Epoch 1/4
1/1 [==============================] - 0s 200ms/step - loss: 24.1234
Epoch 2/4
1/1 [==============================] - 0s 10ms/step - loss: 12.5678
Epoch 3/4
1/1 [==============================] - 0s 10ms/step - loss: 6.5432
Epoch 4/4
1/1 [==============================] - 0s 10ms/step - loss: 3.4567
Смотрите также
-
класс
ModelCheckpoint,
который сохраняет модель во время обучения -
класс
EarlyStopping,
который останавливает обучение при отсутствии улучшений -
класс
CSVLogger,
который записывает метрики обучения в CSV-файл -
класс
TensorBoard,
который визуализирует процесс обучения