Класс Checkpoint
Класс Checkpoint применяется для сохранения
и восстановления состояния объектов TensorFlow.
Он отслеживает переменные моделей, оптимизаторов
и других объектов, позволяя сохранять их на диск
и загружать обратно. Первым параметром передается
словарь отслеживаемых объектов или сами объекты,
вторым - необязательный префикс для файлов
контрольных точек.
Синтаксис
tf.train.Checkpoint(**kwargs)
Пример
Давайте создадим простую модель и оптимизатор,
а затем обернем их в объект Checkpoint:
import tensorflow as tf
model = tf.keras.Sequential([
tf.keras.layers.Dense(2, input_shape=(3,))
])
optimizer = tf.keras.optimizers.Adam()
ckpt = tf.train.Checkpoint(model=model, optimizer=optimizer)
print(ckpt)
Результат выполнения кода:
<tensorflow.python.training.tracking.util.Checkpoint object at 0x7f8b1c0a0d30>
Пример
Давайте сохраним контрольную точку с указанным префиксом и проверим созданные файлы:
import tensorflow as tf
import os
model = tf.keras.Sequential([
tf.keras.layers.Dense(2, input_shape=(3,))
])
optimizer = tf.keras.optimizers.Adam()
ckpt = tf.train.Checkpoint(model=model, optimizer=optimizer)
save_path = ckpt.save('/tmp/ckpt')
print(save_path)
Результат выполнения кода:
"/tmp/ckpt-1"
Пример
Давайте восстановим состояние модели из сохраненной контрольной точки:
import tensorflow as tf
model = tf.keras.Sequential([
tf.keras.layers.Dense(2, input_shape=(3,))
])
optimizer = tf.keras.optimizers.Adam()
ckpt = tf.train.Checkpoint(model=model, optimizer=optimizer)
save_path = ckpt.save('/tmp/ckpt')
model2 = tf.keras.Sequential([
tf.keras.layers.Dense(2, input_shape=(3,))
])
optimizer2 = tf.keras.optimizers.Adam()
ckpt2 = tf.train.Checkpoint(model=model2, optimizer=optimizer2)
status = ckpt2.restore(save_path)
status.expect_partial()
print("Restored")
Результат выполнения кода:
"Restored"
Смотрите также
-
класс
Checkpoint,
который управляет контрольными точками -
метод
save,
который сохраняет контрольную точку на диск -
метод
restore,
который восстанавливает состояние из контрольной точки -
метод
read,
который читает контрольную точку без восстановления