Метод save класса Checkpoint
Метод save класса Checkpoint сохраняет текущие значения отслеживаемых объектов в контрольную точку. Первым параметром метод принимает путь к файлу или каталогу, куда будет сохранено состояние. Вторым параметром можно передать сессию для сохранения. Метод возвращает путь, по которому была сохранена контрольная точка. Класс Checkpoint позволяет сохранять и восстанавливать состояние моделей, оптимизаторов и других объектов TensorFlow.
Синтаксис
tf.train.Checkpoint.save(file_prefix, [session])
Пример
Давайте создадим простую модель и сохраним ее состояние в контрольную точку:
import tensorflow as tf
model = tf.keras.Sequential([
tf.keras.layers.Dense(5, input_shape=(3,))
])
optimizer = tf.keras.optimizers.Adam()
checkpoint = tf.train.Checkpoint(
model=model,
optimizer=optimizer
)
save_path = checkpoint.save('./ckpt/test')
print(save_path)
Результат выполнения кода:
"./ckpt/test-1"
Пример
Давайте сохраним состояние модели после обучения на данных:
import tensorflow as tf
tf.random.set_seed(0)
model = tf.keras.Sequential([
tf.keras.layers.Dense(2, input_shape=(3,))
])
optimizer = tf.keras.optimizers.SGD(learning_rate=0.01)
checkpoint = tf.train.Checkpoint(
model=model,
optimizer=optimizer
)
x = tf.constant([[1, 2, 3], [4, 5, 6]])
y = tf.constant([[1, 0], [0, 1]])
with tf.GradientTape() as tape:
predictions = model(x)
loss = tf.keras.losses.mean_squared_error(y, predictions)
gradients = tape.gradient(loss, model.trainable_variables)
optimizer.apply_gradients(zip(gradients, model.trainable_variables))
save_path = checkpoint.save('./ckpt/trained')
print(save_path)
Результат выполнения кода:
"./ckpt/trained-1"
Пример
Давайте сохраним несколько контрольных точек с разными префиксами:
import tensorflow as tf
model = tf.keras.Sequential([
tf.keras.layers.Dense(3, input_shape=(2,))
])
optimizer = tf.keras.optimizers.Adam()
checkpoint = tf.train.Checkpoint(
model=model,
optimizer=optimizer
)
path1 = checkpoint.save('./ckpt/step1')
path2 = checkpoint.save('./ckpt/step2')
print(path1)
print(path2)
Результат выполнения кода:
"./ckpt/step1-1"
"./ckpt/step2-1"
Смотрите также
-
класс
Checkpoint,
который управляет сохранением и восстановлением состояния -
метод
restore,
который восстанавливает состояние из контрольной точки -
метод
read,
который читает состояние из контрольной точки без восстановления