tf.train.Checkpoint
Управляет сохранением/восстановлением отслеживаемых значений на диск.
tf.train.Checkpoint(
root=None, **kwargs
)
Используется в блокнотах
| Используется в руководстве | Используется в учебных пособиях |
|---|---|
Объекты TensorFlow могут содержать отслеживаемое состояние, такое как tf.Variable, реализации tf.keras.optimizers.Optimizer, итераторы tf.data.Dataset, реализации tf.keras.Layer или реализации tf.keras.Model. Они называются отслеживаемыми объектами.
Объект Checkpoint может быть создан для сохранения одного или группы отслеживаемых объектов в файле контрольной точки. Он поддерживает счётчик для нумерации контрольных точек.
Пример:
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() записывают и считывают контрольные точки на основе объектов, в отличие от 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().__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.
Считывает эту 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 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
)
Восстанавливает обучающую контрольную точку.
Восстанавливает эту Checkpoint и все объекты, от которых она зависит.
Этот метод предназначен для загрузки контрольных точек, созданных с помощью 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() на загруженном состоянии, так как SavedModel, сохранённые из 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. |
| Возвращаемое значение | |
|---|---|
| Полный путь к контрольной точке. |
sync
sync()
Ожидание завершения всех операций сохранения или восстановления.
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/api_docs/python/tf/train/Checkpoint