Spec-Zone.ru › TensorFlow 2.9

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.
Возвращаемые значения
Объект состояния загрузки, который можно использовать для утверждений о состоянии восстановления контрольной точки.

Возвращаемый объект состояния имеет следующие методы:

  • assert_consumed(): Вызывает исключение, если какие-либо переменные не сопоставлены: либо проверенные значения, у которых нет соответствующего Python-объекта, или Python-объекты в графе зависимостей без значений в контрольной точке. Этот метод возвращает объект состояния и, следовательно, может быть объединён с другими утверждениями.

  • assert_existing_objects_matched(): Вызывает исключение, если какие-либо существующие Python-объекты в графе зависимостей не сопоставлены. В отличие от assert_consumed, это утверждение пройдёт, если значения в контрольной точке не имеют соответствующих Python-объектов. Например, объект tf.keras.Layer, который ещё не был построен и, следовательно, не создал никаких переменных, пройдёт это утверждение, но не пройдёт assert_consumed. Полезно при загрузке части большей контрольной точки в новую Python-программу, например, контрольной точки обучения с tf.compat.v1.train.Optimizer была сохранена, но загружается только состояние, необходимое для вывода. Этот метод возвращает объект состояния и, следовательно, может быть объединён с другими утверждениями.

  • assert_nontrivial_match(): Утверждает, что что-то помимо корневого объекта было сопоставлено. Это очень слабое утверждение, но полезно для проверки корректности в коде библиотеки, где объекты могут существовать в контрольной точке, которые не были созданы в Python, и некоторые Python-объекты могут не иметь сохранённого значения.

  • expect_partial(): Заглушает предупреждения о неполных восстановлениях контрольных точек. В противном случае предупреждения выводятся для неиспользованных частей файла контрольной точки или объекта, когда объект Checkpoint удаляется (часто при завершении программы).

Исключения
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

Spec-Zone.ru

Настройки Оффлайн Что нового Помощь О нас
Spec-Zone .ru
спецификации, руководства, описания, API