Spec-Zone.ru › TensorFlow 2.9

tf.compat.v1.train.Checkpoint

Группирует отслеживаемые объекты, сохраняя и восстанавливая их.

tf.compat.v1.train.Checkpoint(
    **kwargs
)

Конструктор 's принимает именованные аргументы, значениями которых являются типы, содержащие отслеживаемое состояние, такие как реализации tf.compat.v1.train.Optimizer, tf.Variable, реализации tf.keras.Layer или реализации tf.keras.Model. Он сохраняет эти значения с контрольной точкой и поддерживает save_counter для нумерации контрольных точек.

Пример использования при построении графа:

import tensorflow as tf
import os

checkpoint_directory = "/tmp/training_checkpoints"
checkpoint_prefix = os.path.join(checkpoint_directory, "ckpt")

checkpoint = tf.train.Checkpoint(optimizer=optimizer, model=model)
status = checkpoint.restore(tf.train.latest_checkpoint(checkpoint_directory))
train_op = optimizer.minimize( ... )
status.assert_consumed()  # Optional sanity checks.
with tf.compat.v1.Session() as session:
  # Use the Session to restore variables, or initialize them if
  # tf.train.latest_checkpoint returned None.
  status.initialize_or_restore(session)
  for _ in range(num_training_steps):
    session.run(train_op)
  checkpoint.save(file_prefix=checkpoint_prefix)

Пример использования с включённым выполнением в режиме eager:

import tensorflow as tf
import os

tf.compat.v1.enable_eager_execution()

checkpoint_directory = "/tmp/training_checkpoints"
checkpoint_prefix = os.path.join(checkpoint_directory, "ckpt")

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

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

Когда переменные присваиваются нескольким рабочим процессам, каждый рабочий процесс записывает свою секцию контрольной точки. Затем эти секции объединяются/переиндексируются, чтобы работать как единая контрольная точка. Это позволяет избежать копирования всех переменных одному рабочему процессу, но требует, чтобы все рабочие процессы имели доступ к общему файловой системе.

Хотя tf.keras.Model.save_weights и tf.train.Checkpoint.save сохраняют в одном формате, обратите внимание, что корень получившейся контрольной точки — это объект, к которому прикреплён метод save. Это означает, что сохранение tf.keras.Model с помощью save_weights и загрузка в tf.train.Checkpoint с присоединённым Model (или наоборот) не будут соответствовать переменным Model. Подробности см. в руководстве по тренировочным контрольным точкам. Для тренировочных контрольных точек предпочтительно использовать tf.train.Checkpoint вместо tf.keras.Model.save_weights.

Args
**kwargs Именованные аргументы устанавливаются в качестве атрибутов этого объекта и сохраняются с контрольной точкой. Значения должны быть отслеживаемыми объектами.
Raises
ValueError Если объекты в kwargs не являются отслеживаемыми.
Attributes
save_counter Увеличивается при вызове save(). Используется для нумерации контрольных точек.

Методы

restore

Просмотреть исходный код

restore(
    save_path
)

Восстановление тренировочной контрольной точки.

Восстанавливает эту Checkpoint и любые зависящие от неё объекты.

При выполнении в режиме eager, либо сразу присваивает значения, если переменные для восстановления уже созданы, либо откладывает восстановление до создания переменных. Зависимости, добавленные после этого вызова, будут сопоставлены, если в контрольной точке есть соответствующий объект (запрос на восстановление будет помещён в очередь в любом отслеживаемом объекте, ожидающем добавления ожидаемой зависимости).

При построении графа операции восстановления добавляются в граф, но не выполняются немедленно.

checkpoint = tf.train.Checkpoint( ... )
checkpoint.restore(path)

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

checkpoint = tf.train.Checkpoint( ... )
checkpoint.restore(path).assert_consumed()

При построении графа assert_consumed() указывает, что все операции восстановления, которые будут созданы для этой контрольной точки, были созданы. Их можно выполнить с помощью метода run_restore_ops() объекта состояния:

checkpoint.restore(path).assert_consumed().run_restore_ops()

Если контрольная точка не была полностью использована, список операций восстановления будет расти по мере добавления большего количества объектов в граф зависимостей.

Чтобы убедиться, что все переменные в объекте Python восстановили значения из контрольной точки, используйте assert_existing_objects_matched(). Эта проверка полезна, когда она вызывается после того, как переменные в вашем графе были созданы.

Контрольные точки на основе имён tf.compat.v1.train.Saver могут быть загружены с помощью этого метода. Имена используются для сопоставления переменных. Никаких операций восстановления не создаётся/выполняется до тех пор, пока run_restore_ops() или initialize_or_restore() не будут вызваны для возвращаемого объекта состояния при построении графа, но восстановление при создании имеет место при выполнении в режиме eager. Перекодируйте контрольные точки на основе имён с помощью tf.train.Checkpoint.save как можно скорее.

Args
save_path Путь к контрольной точке, возвращаемый методом save или tf.train.latest_checkpoint. Если None (как в случае отсутствия последней контрольной точки для возврата tf.train.latest_checkpoint), возвращает объект, который может запускать инициализаторы объектов в графе зависимостей. Если контрольная точка была записана с помощью tf.compat.v1.train.Saver на основе имён, имена используются для сопоставления переменных.
Возвращает
Объект состояния загрузки, который может использоваться для проверки состояния восстановления контрольной точки и запуска операций инициализации/восстановления.

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

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

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

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

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

  • initialize_or_restore(session=None): При построении графа, выполняет инициализаторы переменных, если save_path равно None, а иначе выполняет операции восстановления. Если session не указан явно, используется стандартная сессия. Нет эффекта при выполнении в режиме eager (переменные инициализируются или восстанавливаются в режиме eager).

  • run_restore_ops(session=None): При построении графа, выполняет операции восстановления. Если session не указан явно, используется стандартная сессия. Нет эффекта при выполнении в режиме eager (операции восстановления выполняются в режиме eager). Может быть вызван только тогда, когда save_path не равно None.

save

Просмотреть исходный код

save(
    file_prefix, session=None
)

Сохраняет тренировочную контрольную точку и предоставляет базовое управление контрольными точками.

Сохранённая контрольная точка включает переменные, созданные этим объектом, и любые отслеживаемые объекты, от которых он зависит, на момент вызова Checkpoint.save().

save — это базовая обёртка над методом write, последовательно нумерующая контрольные точки с помощью save_counter и обновляющая метаданные, используемые tf.train.latest_checkpoint. Более продвинутое управление контрольными точками, например, сбор мусора и пользовательская нумерация, могут быть предоставлены другими утилитами, которые также обертывают write (tf.train.CheckpointManager, например).

Аргументы
file_prefix Префикс для имён файлов контрольных точек (/путь/к/каталогу/и_префикс). Имена генерируются на основе этого префикса и Checkpoint.save_counter.
session Сессия для оценки переменных. Игнорируется при выполнении в режиме eager. Если не указана при построении графа, используется по умолчанию.
Возвращаемое значение
Полный путь к контрольной точке.

write

Просмотреть исходный код

write(
    file_prefix, session=None
)

Записывает контрольную точку обучения.

Контрольная точка включает переменные, созданные этим объектом, и любые отслеживаемые объекты, от которых он зависит на момент вызова Checkpoint.write().

write не нумерует контрольные точки, не инкрементирует save_counter, и не обновляет метаданные, используемые tf.train.latest_checkpoint. Она в первую очередь предназначена для использования утилитами управления контрольными точками более высокого уровня. save предоставляет очень базовую реализацию этих функций.

Аргументы
file_prefix Префикс для имён файлов контрольных точек (/путь/к/каталогу/и_префикс).
session Сессия для оценки переменных. Игнорируется при выполнении в режиме eager. Если не указана при построении графа, используется по умолчанию.
Возвращаемое значение
Полный путь к контрольной точке (т.е. 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/compat/v1/train/Checkpoint

Spec-Zone.ru

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