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

Метод restore класса Checkpoint

Метод restore класса Checkpoint восстанавливает значения отслеживаемых объектов из сохраненного чекпоинта. Первым параметром метод принимает путь к чекпоинту, вторым - необязательный объект сессии.

Синтаксис

checkpoint.restore(save_path)

Пример

Давайте создадим чекпоинт с переменной, сохраним ее, а затем восстановим в новую переменную:

import tensorflow as tf # create and save a variable v1 = tf.Variable([1, 2, 3, 4, 5], name='v1') ckpt = tf.train.Checkpoint(v1=v1) save_path = ckpt.save('./model.ckpt') print("saved:", save_path) # restore into a new variable v2 = tf.Variable([0, 0, 0, 0, 0], name='v2') ckpt2 = tf.train.Checkpoint(v1=v2) ckpt2.restore(save_path) print(v2)

Результат выполнения кода:

saved: ./model.ckpt-1 <tf.Variable 'v2:0' shape=(5,) dtype=int32, numpy=array([1, 2, 3, 4, 5], dtype=int32)>

Пример

Давайте восстановим состояние нескольких переменных из одного чекпоинта:

import tensorflow as tf # create and save two variables a = tf.Variable([1, 2, 3], name='a') b = tf.Variable([4, 5, 6], name='b') ckpt = tf.train.Checkpoint(a=a, b=b) save_path = ckpt.save('./model.ckpt') print("saved:", save_path) # restore into new variables a2 = tf.Variable([0, 0, 0], name='a2') b2 = tf.Variable([0, 0, 0], name='b2') ckpt2 = tf.train.Checkpoint(a=a2, b=b2) ckpt2.restore(save_path) print(a2) print(b2)

Результат выполнения кода:

saved: ./model.ckpt-1 <tf.Variable 'a2:0' shape=(3,) dtype=int32, numpy=array([1, 2, 3], dtype=int32)> <tf.Variable 'b2:0' shape=(3,) dtype=int32, numpy=array([4, 5, 6], dtype=int32)>

Пример

Давайте восстановим переменные внутри модели Keras:

import tensorflow as tf # create and save a model model = tf.keras.Sequential([ tf.keras.layers.Dense(2, input_shape=(3,)) ]) model.compile(optimizer='sgd', loss='mse') ckpt = tf.train.Checkpoint(model=model) save_path = ckpt.save('./model.ckpt') print("saved:", save_path) # create a new model and restore weights model2 = tf.keras.Sequential([ tf.keras.layers.Dense(2, input_shape=(3,)) ]) model2.compile(optimizer='sgd', loss='mse') ckpt2 = tf.train.Checkpoint(model=model2) ckpt2.restore(save_path) print(model2.weights[0])

Результат выполнения кода:

saved: ./model.ckpt-1 <tf.Variable 'dense_1/kernel:0' shape=(3, 2) dtype=float32, numpy= array([[...]], dtype=float32)>

Смотрите также

  • класс Checkpoint,
    который управляет сохранением и восстановлением состояния
  • метод save,
    который сохраняет состояние в чекпоинт
  • метод read,
    который читает значения из чекпоинта
  • метод restore,
    который восстанавливает значения из чекпоинта
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить