tf.train.Checkpoint
| Просмотр исходного кода на GitHub |
Группирует отслеживаемые объекты, сохраняя и восстанавливая их.
tf.train.Checkpoint(
**kwargs
)
Конструктор класса Checkpoint принимает именованные аргументы, значениями которых являются типы, содержащие отслеживаемое состояние, такие как реализации tf.keras.optimizers.Optimizer, объекты tf.Variable, итераторы tf.data.Dataset, реализации tf.keras.Layer или реализации tf.keras.Model. Он сохраняет эти значения вместе с контрольной точкой и поддерживает счётчик для нумерации контрольных точек.
Пример использования:
import tensorflow as tf import os checkpoint_directory = "/tmp/training_checkpoints" checkpoint_prefix = os.path.join(checkpoint_directory, "ckpt") # Create a Checkpoint that will manage two objects with trackable state, # one we name "optimizer" and the other we name "model". checkpoint = tf.train.Checkpoint(optimizer=optimizer, model=model) status = checkpoint.restore(tf.train.latest_checkpoint(checkpoint_directory)) for _ in range(num_training_steps): optimizer.minimize( ... ) # Variables will be restored on creation. status.assert_consumed() # Optional sanity checks. checkpoint.save(file_prefix=checkpoint_prefix)
Checkpoint.save() и Checkpoint.restore() записывают и считывают контрольные точки на основе объектов, в отличие от tf.compat.v1.train.Saver TensorFlow 1.x, который записывает и считывает контрольные точки на основе variable.name. Сохранение контрольных точек на основе объектов сохраняет граф зависимостей между объектами Python (Layerы, Optimizerы, Variableы и т. д.) с именованными рёбрами, и этот граф используется для сопоставления переменных при восстановлении контрольной точки. Это может быть более устойчиво к изменениям в программе Python и помогает поддерживать восстановление при создании переменных.
Объекты Checkpoint зависят от объектов, переданных в качестве именованных аргументов их конструкторам, и каждая зависимость получает имя, совпадающее с именем именованного аргумента, для которого она была создана. Классы TensorFlow, такие как Layer и Optimizer автоматически добавляют зависимости на свои собственные переменные (например, "kernel" и "bias" для tf.keras.layers.Dense). Наследование от tf.keras.Model упрощает управление зависимостями в пользовательских классах, поскольку Model реагирует на присвоение атрибутов. Например:
class Regress(tf.keras.Model):
def __init__(self):
super(Regress, self).__init__()
self.input_transform = tf.keras.layers.Dense(10)
# ...
def call(self, inputs):
x = self.input_transform(inputs)
# ...
Этот Model имеет зависимость с именем "input_transform" от слоя Dense, который, в свою очередь, зависит от своих переменных. В результате сохранение экземпляра Regress с помощью tf.train.Checkpoint также сохранит все переменные, созданные слоем Dense.
Когда переменные присваиваются нескольким рабочим процессам, каждый рабочий процесс записывает свою часть контрольной точки. Затем эти части объединяются/переиндексируются, чтобы вести себя как единая контрольная точка. Это позволяет избежать копирования всех переменных в один рабочий процесс, но требует, чтобы все рабочие процессы имели доступ к одному файловому хранилищу.
Хотя tf.keras.Model.save_weights и tf.train.Checkpoint.save сохраняют в одном формате, обратите внимание, что корень полученной контрольной точки — это объект, к которому прикреплён метод сохранения. Это означает, что сохранение tf.keras.Model с помощью save_weights и загрузка в tf.train.Checkpoint с прикреплённым Model (или наоборот) не приведет к совпадению переменных Model. Подробности см. в руководстве по контрольным точкам обучения. Предпочтительнее использовать tf.train.Checkpoint вместо tf.keras.Model.save_weights для контрольных точек обучения.
| Аргументы | |
|---|---|
**kwargs | Именованные аргументы устанавливаются в качестве атрибутов этого объекта и сохраняются вместе с контрольной точкой. Значения должны быть отслеживаемыми объектами. |
| Исключения | |
|---|---|
ValueError | Если объекты в kwargs не отслеживаются. |
| Атрибуты | |
|---|---|
save_counter | Увеличивается при вызове save(). Используется для нумерации контрольных точек. |
Методы
read
read(
save_path, options=None
)
Чтение контрольной точки обучения, записанной с помощью write.
Читает эту Checkpoint и все объекты, от которых она зависит.
Этот метод аналогичен restore() но не ожидает переменную save_counter в контрольной точке. Он восстанавливает только те объекты, от которых контрольная точка уже зависит.
Этот метод предназначен в основном для использования в утилитах управления контрольными точками более высокого уровня, которые используют write() вместо save() и имеют свои собственные механизмы для нумерации и отслеживания контрольных точек.
Пример использования:
# Create a checkpoint with write()
ckpt = tf.train.Checkpoint(v=tf.Variable(1.))
path = ckpt.write('/tmp/my_checkpoint')
# Later, load the checkpoint with read()
# With restore() assert_consumed() would have failed.
checkpoint.read(path).assert_consumed()
# You can also pass options to restore(). For example this
# runs the IO ops on the localhost:
options = tf.CheckpointOptions(experimental_io_device="/job:localhost")
checkpoint.read(path, options=options)
| Аргументы | |
|---|---|
save_path | Путь к контрольной точке, возвращаемый write . |
options | Необязательный объект tf.train.CheckpointOptions. |
| Возвращаемое значение | |
|---|---|
Объект состояния загрузки, который можно использовать для утверждений о состоянии восстановления контрольной точки. Подробности см. в restore . |
restore
restore(
save_path, options=None
)
Восстановление контрольной точки обучения.
Восстанавливает эту Checkpoint и все объекты, от которых она зависит.
Этот метод предназначен для загрузки контрольных точек, созданных save(). Для контрольных точек, созданных write(), используйте метод read(), который не ожидает переменной save_counter, добавленной save().
restore() либо сразу присваивает значения, если переменные для восстановления уже созданы, либо откладывает восстановление до создания переменных. Зависимости, добавленные после этого вызова, будут сопоставляться, если в контрольной точке есть соответствующий объект (запрос на восстановление будет помещен в очередь в любом отслеживаемом объекте, ожидающем добавления ожидаемой зависимости).
Для обеспечения завершения загрузки и отсутствия дальнейших присвоений используйте метод assert_consumed() объекта состояния, возвращаемого restore():
checkpoint = tf.train.Checkpoint( ... ) checkpoint.restore(path).assert_consumed() # You can additionally pass options to restore(): options = tf.CheckpointOptions(experimental_io_device="/job:localhost") checkpoint.restore(path, options=options).assert_consumed()
Будет выброшено исключение, если какие-либо объекты Python в графе зависимостей не были найдены в контрольной точке или если какие-либо сохранённые значения не имеют соответствующего объекта Python.
Можно загрузить контрольные точки на основе имён tf.compat.v1.train.Saver из TensorFlow 1.x с помощью этого метода. Имена используются для сопоставления переменных. Перекодируйте контрольные точки на основе имён с помощью tf.train.Checkpoint.save как можно скорее.
| Аргументы | |
|---|---|
save_path | Путь к контрольной точке, возвращаемый save или tf.train.latest_checkpoint. Если контрольная точка была записана с помощью tf.compat.v1.train.Saver на основе имён, имена используются для сопоставления переменных. |
options | Необязательный объект tf.train.CheckpointOptions. |
| Возвращаемое значение | |
|---|---|
| Объект состояния загрузки, который можно использовать для утверждений о состоянии восстановления контрольной точки. Возвращаемый объект состояния имеет следующие методы:
|
save
save(
file_prefix, options=None
)
Сохраняет контрольную точку обучения и предоставляет базовое управление контрольными точками.
Сохраненная контрольная точка включает переменные, созданные этим объектом, и любые отслеживаемые объекты, от которых он зависит, на момент вызова Checkpoint.save().
save — это базовая утилита для обёртки метода write, которая последовательно нумерует контрольные точки с помощью save_counter и обновляет метаданные, используемые tf.train.latest_checkpoint. Более продвинутое управление контрольными точками, например, сборка мусора и настройка нумерации, могут быть предоставлены другими утилитами, которые также оборачивают write и read (например, tf.train.CheckpointManager).
step = tf.Variable(0, name="step")
checkpoint = tf.Checkpoint(step=step)
checkpoint.save("/tmp/ckpt")
# Later, read the checkpoint with restore()
checkpoint.restore("/tmp/ckpt").assert_consumed()
# You can also pass options to save() and restore(). For example this
# runs the IO ops on the localhost:
options = tf.CheckpointOptions(experimental_io_device="/job:localhost")
checkpoint.save("/tmp/ckpt", options=options)
# Later, read the checkpoint with restore()
checkpoint.restore("/tmp/ckpt", options=options).assert_consumed()
| Аргументы | |
|---|---|
file_prefix | Префикс для имён файлов контрольных точек (/path/to/directory/and_a_prefix). Имена генерируются на основе этого префикса и Checkpoint.save_counter. |
options | Необязательный объект tf.train.CheckpointOptions. |
| Возвращаемое значение | |
|---|---|
| Полный путь к контрольной точке. |
write
write(
file_prefix, options=None
)
Записывает контрольную точку обучения.
Контрольная точка включает переменные, созданные этим объектом, и любые отслеживаемые объекты, от которых он зависит, на момент вызова Checkpoint.write().
write не нумерует контрольные точки, не инкрементирует save_counter, и не обновляет метаданные, используемые tf.train.latest_checkpoint. В основном он предназначен для использования в утилитах управления контрольными точками более высокого уровня. save предоставляет очень базовую реализацию этих функций.
Контрольные точки, записанные с помощью write , должны читаться с помощью read.
Пример использования:
step = tf.Variable(0, name="step")
checkpoint = tf.Checkpoint(step=step)
checkpoint.write("/tmp/ckpt")
# Later, read the checkpoint with read()
checkpoint.read("/tmp/ckpt").assert_consumed()
# You can also pass options to write() and read(). For example this
# runs the IO ops on the localhost:
options = tf.CheckpointOptions(experimental_io_device="/job:localhost")
checkpoint.write("/tmp/ckpt", options=options)
# Later, read the checkpoint with read()
checkpoint.read("/tmp/ckpt", options=options).assert_consumed()
| Аргументы | |
|---|---|
file_prefix | Префикс для имён файлов контрольных точек (/path/to/directory/and_a_prefix). |
options | Необязательный объект tf.train.CheckpointOptions. |
| Возвращаемое значение | |
|---|---|
Полный путь к контрольной точке (т.е. file_prefix). |
© 2020 The TensorFlow Authors. All rights reserved.
Licensed under the Creative Commons Attribution License 3.0.
Code samples licensed under the Apache 2.0 License.
https://www.tensorflow.org/versions/r2.3/api_docs/python/tf/train/Checkpoint