Метод save класса CheckpointManager
Метод save класса CheckpointManager
сохраняет текущее состояние модели в контрольную
точку. Первым параметром метод принимает номер
шага, вторым - необязательный словарь
дополнительных метрик. Метод автоматически
нумерует чекпоинты, удаляет старые и обновляет
атрибут latest_checkpoint.
Синтаксис
CheckpointManager.save(checkpoint_number, [options])
Пример
Давайте создадим объект CheckpointManager
и сохраним одну контрольную точку:
import tensorflow as tf
model = tf.keras.Sequential([
tf.keras.layers.Dense(1)
])
model.compile(optimizer='sgd', loss='mse')
ckpt = tf.train.Checkpoint(model=model)
manager = tf.train.CheckpointManager(
ckpt, './ckpt', max_to_keep=3
)
path = manager.save(checkpoint_number=1)
print(path)
print(manager.latest_checkpoint)
Результат выполнения кода:
"./ckpt/ckpt-1"
"./ckpt/ckpt-1"
Пример
Давайте сохраним несколько контрольных точек
и посмотрим, как работает ограничение
max_to_keep:
import tensorflow as tf
model = tf.keras.Sequential([
tf.keras.layers.Dense(1)
])
model.compile(optimizer='sgd', loss='mse')
ckpt = tf.train.Checkpoint(model=model)
manager = tf.train.CheckpointManager(
ckpt, './ckpt', max_to_keep=2
)
manager.save(checkpoint_number=1)
manager.save(checkpoint_number=2)
manager.save(checkpoint_number=3)
print(manager.checkpoints)
print(manager.latest_checkpoint)
Результат выполнения кода:
['./ckpt/ckpt-2', './ckpt/ckpt-3']
"./ckpt/ckpt-3"
Пример
Давайте передадим дополнительные метрики при сохранении контрольной точки:
import tensorflow as tf
model = tf.keras.Sequential([
tf.keras.layers.Dense(1)
])
model.compile(optimizer='sgd', loss='mse')
ckpt = tf.train.Checkpoint(model=model)
manager = tf.train.CheckpointManager(
ckpt, './ckpt', max_to_keep=3
)
path = manager.save(
checkpoint_number=1,
options=tf.train.CheckpointOptions(
experimental_io_device='/job:localhost'
)
)
print(path)
Результат выполнения кода:
"./ckpt/ckpt-1"
Смотрите также
-
класс
CheckpointManager,
который управляет сохранением контрольных точек -
метод
restore_or_initialize,
который восстанавливает или инициализирует модель -
атрибут
latest_checkpoint,
который хранит путь к последней контрольной точке -
атрибут
checkpoints,
который содержит список всех контрольных точек