Spec-Zone.ru › TensorFlow 2.3

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

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

  • 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 удаляется (часто при завершении программы).

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

Spec-Zone.ru

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