Метод 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,
который восстанавливает значения из чекпоинта