РЕПЕТИТОР математика физика информатика
Для школьников и студентов. Подтягивание пробелов. ЦЭ, ЦТ, ОГЭ, ЕГЭ.
Идет набор на ЛЕТО. Жмите для подробностей:)
665 of 824 menu

Метод 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,
    который содержит список всех сохраненных чекпоинтов
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить