Метод restore_or_initialize
Метод restore_or_initialize класса CheckpointManager
проверяет наличие сохраненных чекпоинтов и автоматически
восстанавливает последний из них. Если чекпоинтов не найдено,
метод выполняет инициализацию переменных модели. Это удобный
способ продолжить обучение с сохраненного состояния или начать
обучение с нуля без дополнительных проверок. Метод не принимает
обязательных параметров и возвращает объект, содержащий статус
восстановления.
Синтаксис
status = checkpoint_manager.restore_or_initialize()
Пример
Давайте создадим простую модель, менеджер чекпоинтов и вызовем
метод restore_or_initialize, когда чекпоинтов еще нет:
import tensorflow as tf
model = tf.keras.Sequential([
tf.keras.layers.Dense(2, input_shape=(3,))
])
optimizer = tf.keras.optimizers.Adam()
checkpoint = tf.train.Checkpoint(
model=model,
optimizer=optimizer
)
manager = tf.train.CheckpointManager(
checkpoint,
directory='./ckpt',
max_to_keep=3
)
status = manager.restore_or_initialize()
print(status)
Результат выполнения кода:
<tf.train.CheckpointManager.RestoreOrInitializeStatus object at 0x...>
Пример
Давайте сохраним чекпоинт через метод save, а затем
восстановим его с помощью restore_or_initialize:
import tensorflow as tf
model = tf.keras.Sequential([
tf.keras.layers.Dense(2, input_shape=(3,))
])
optimizer = tf.keras.optimizers.Adam()
checkpoint = tf.train.Checkpoint(
model=model,
optimizer=optimizer
)
manager = tf.train.CheckpointManager(
checkpoint,
directory='./ckpt',
max_to_keep=3
)
save_path = manager.save()
print('Saved checkpoint:', save_path)
status = manager.restore_or_initialize()
print('Restored:', status)
Результат выполнения кода:
"Saved checkpoint: ./ckpt/ckpt-1"
"Restored: <tf.train.CheckpointManager.RestoreOrInitializeStatus object at 0x...>"
Пример
Давайте проверим, что после восстановления веса модели совпадают с сохраненными значениями:
import tensorflow as tf
model = tf.keras.Sequential([
tf.keras.layers.Dense(2, input_shape=(3,))
])
optimizer = tf.keras.optimizers.Adam()
checkpoint = tf.train.Checkpoint(
model=model,
optimizer=optimizer
)
manager = tf.train.CheckpointManager(
checkpoint,
directory='./ckpt',
max_to_keep=3
)
model(tf.constant([[1.0, 2.0, 3.0]]))
weights_before = model.get_weights()[0].copy()
manager.save()
model.set_weights([w * 0 for w in model.get_weights()])
manager.restore_or_initialize()
weights_after = model.get_weights()[0]
print('Weights match:', (weights_before == weights_after).all())
Результат выполнения кода:
"Weights match: True"
Смотрите также
-
класс
CheckpointManager,
который управляет сохранением и восстановлением чекпоинтов -
метод
save,
который сохраняет новый чекпоинт -
атрибут
latest_checkpoint,
который хранит путь к последнему чекпоинту -
атрибут
checkpoints,
который содержит список всех сохраненных чекпоинтов