Spec-Zone.ru › TensorFlow 2.4

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() записывают и считывают контрольные точки на основе объектов, в отличие от tf.compat.v1.train.Saver TensorFlow 1.x, который записывает и считывает контрольные точки на основе variable.name. Сохранение контрольных точек на основе объектов сохраняет граф зависимостей между объектами Python (Layers, Optimizers, Variables и т. д.) с именованными ребрами, и этот граф используется для сопоставления переменных при восстановлении контрольной точки. Это может быть более устойчивым к изменениям в программе Python и помогает поддерживать восстановление при создании переменных.

Объекты Checkpoint зависят от объектов, переданных в качестве ключевых аргументов в их конструкторы, и каждая зависимость получает имя, идентичное имени ключевого аргумента, для которого она была создана. Классы TensorFlow, такие как Layer и Optimizers, автоматически добавляют зависимости от своих собственных переменных (например, «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 Корневой объект для контрольной точки.
**kwargs Ключевые аргументы устанавливаются в качестве атрибутов этого объекта и сохраняются вместе с контрольной точкой. Значения должны быть отслеживаемыми объектами.
Исключения
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.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 как можно скорее.

Загрузка из контрольных точек 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.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 Префикс для имён файлов контрольных точек (/путь/к/каталогу/и_префикс). Имена генерируются на основе этого префикса и 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 Префикс для имён файлов контрольных точек (/путь/к/каталогу/и_префикс).
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.4/api_docs/python/tf/train/Checkpoint

Spec-Zone.ru

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