tf.train.Checkpoint
| Просмотреть исходный код на GitHub |
Управляет сохранением/восстановлением отслеживаемых значений на диск.
tf.train.Checkpoint(
root=None, **kwargs
)
Объекты TensorFlow могут содержать отслеживаемое состояние, например, tf.Variable, реализации tf.keras.optimizers.Optimizer, итераторы tf.data.Dataset, реализации tf.keras.Layer или реализации tf.keras.Model. Они называются отслеживаемыми объектами.
Объект Checkpoint может быть создан для сохранения одного или группы отслеживаемых объектов в файле контрольной точки. Он поддерживает save_counter для нумерации контрольных точек.
Пример:
model = tf.keras.Model(...)
checkpoint = tf.train.Checkpoint(model)
# Save a checkpoint to /tmp/training_checkpoints-{save_counter}. Every time
# checkpoint.save is called, the save counter is increased.
save_path = checkpoint.save('/tmp/training_checkpoints')
# Restore the checkpointed values to the `model` object.
checkpoint.restore(save_path)
Пример 2:
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() записывают и считывают контрольные точки на основе объектов, в отличие от TensorFlow 1.x's tf.compat.v1.train.Saver, которые записывают и считывают контрольные точки на основе variable.name. Контрольные точки на основе объектов сохраняют граф зависимостей между Python-объектами (Layers, Optimizers, Variables и т.д.) с именованными ребрами, и этот граф используется для сопоставления переменных при восстановлении контрольной точки. Это может быть более устойчиво к изменениям в 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.
Когда переменные назначаются нескольким рабочим процессам, каждый рабочий процесс записывает свою часть контрольной точки. Затем эти части объединяются/переиндексируются, чтобы вести себя как единая контрольная точка. Это предотвращает копирование всех переменных в один рабочий процесс, но требует, чтобы все рабочие процессы видели общую файловую систему.
Эта функция немного отличается от функции Keras Model save_weights . tf.keras.Model.save_weights создаёт файл контрольной точки с именем, указанным в filepath, а tf.train.Checkpoint нумерует контрольные точки, используя filepath в качестве префикса для имён файлов контрольных точек. Помимо этого, model.save_weights() и tf.train.Checkpoint(model).save() эквивалентны.
Подробности см. в руководстве по контрольным точкам обучения.
| Аргументы | |
|---|---|
root | Корневой объект для сохранения контрольной точки. root может быть отслеживаемым объектом или WeakRef отслеживаемого объекта. |
**kwargs | Именованные аргументы устанавливаются в качестве атрибутов этого объекта и сохраняются вместе с контрольной точкой. Все kwargs должны быть отслеживаемыми объектами или вложенной структурой отслеживаемых объектов (list, dict, или tuple). |
| Исключения | |
|---|---|
ValueError | Если root или объекты в kwargs не являются отслеживаемыми. Исключение ValueError также возникает, если объект root отслеживает разные объекты, чем те, которые перечислены в атрибутах в kwargs (например, root.child = A и tf.train.Checkpoint(root, child=B) несовместимы). |
| Атрибуты | |
|---|---|
save_counter | Увеличивается при вызове save(). Используется для нумерации контрольных точек. |
Методы
read
read(
save_path, options=None
)
Считывает контрольную точку обучения, созданную с помощью write.
Считывает эту контрольную точку и все объекты, от которых она зависит.
Этот метод аналогичен методу 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 read(). For example this
# runs the IO ops on the localhost:
options = tf.train.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
)
Восстанавливает контрольную точку обучения.
Восстанавливает эту контрольную точку и все объекты, от которых она зависит.
Этот метод предназначен для использования при загрузке контрольных точек, созданных save(). Для контрольных точек, созданных write(), используйте метод read(), который не ожидает переменной save_counter, добавленной save().
restore() либо сразу присваивает значения, если переменные для восстановления уже созданы, либо откладывает восстановление до создания переменных. Зависимости, добавленные после этого вызова, будут сопоставлены, если у них есть соответствующий объект в контрольной точке (запрос на восстановление будет помещен в очередь в любом отслеживаемом объекте, ожидающем добавления ожидаемой зависимости).
checkpoint = tf.train.Checkpoint( ... ) checkpoint.restore(path) # You can additionally pass options to restore(): options = tf.CheckpointOptions(experimental_io_device="/job:localhost") checkpoint.restore(path, options=options)
Чтобы убедиться, что загрузка завершена и больше не будет происходить отложенного восстановления, используйте метод assert_consumed() объекта состояния, возвращённого restore():
checkpoint.restore(path, options=options).assert_consumed()
Утверждение вызовет ошибку, если какие-либо Python-объекты в графе зависимостей не были найдены в контрольной точке или если какие-либо сохранённые значения не имеют соответствующего Python-объекта.
Контрольные точки на основе имён tf.compat.v1.train.Saver из TensorFlow 1.x можно загрузить с помощью этого метода. Имена используются для сопоставления переменных. Перекодируйте контрольные точки на основе имён с помощью tf.train.Checkpoint.save как можно скорее.
Загрузка из контрольных точек SavedModel
Чтобы загрузить значения из SavedModel, просто передайте каталог SavedModel в checkpoint.restore:
model = tf.keras.Model(...) tf.saved_model.save(model, path) # or model.save(path, save_format='tf') checkpoint = tf.train.Checkpoint(model) checkpoint.restore(path).expect_partial()
В этом примере вызывается expect_partial() на загруженном состоянии, так как SavedModels, сохраненные из Keras, часто генерируют дополнительные ключи в контрольной точке. В противном случае программа выводит много предупреждений об неиспользованных ключах при выходе.
| Аргументы | |
|---|---|
save_path | Путь к контрольной точке, возвращаемый save или tf.train.latest_checkpoint. Если контрольная точка была записана с помощью tf.compat.v1.train.Saver на основе имён, для сопоставления переменных используются имена. Этот путь также может быть каталогом SavedModel. |
options | Необязательный объект tf.train.CheckpointOptions. |
| Возвращаемые значения | |
|---|---|
| Объект состояния загрузки, который можно использовать для утверждений о состоянии восстановления контрольной точки. Возвращаемый объект состояния имеет следующие методы:
|
| Исключения | |
|---|---|
NotFoundError | если контрольная точка или SavedModel не найдены по адресу save_path. |
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.train.Checkpoint(step=step)
checkpoint.save("/tmp/ckpt")
# Later, read the checkpoint with restore()
checkpoint.restore("/tmp/ckpt-1")
# You can also pass options to save() and restore(). For example this
# runs the IO ops on the localhost:
options = tf.train.CheckpointOptions(experimental_io_device="/job:localhost")
checkpoint.save("/tmp/ckpt", options=options)
# Later, read the checkpoint with restore()
checkpoint.restore("/tmp/ckpt-1", options=options)
| Аргументы | |
|---|---|
file_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")
# 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)
| Аргументы | |
|---|---|
file_prefix | Префикс, используемый для имён файлов контрольных точек (/путь/к/каталогу/и_префикс). |
options | Необязательный объект tf.train.CheckpointOptions. |
| Возвращаемое значение | |
|---|---|
Полный путь к контрольной точке (т.е. file_prefix). |
© 2022 The TensorFlow Authors. All rights reserved.
Licensed under the Creative Commons Attribution License 4.0.
Code samples licensed under the Apache 2.0 License.
https://www.tensorflow.org/versions/r2.9/api_docs/python/tf/train/Checkpoint