tf.compat.v1.train.Checkpoint
Группирует отслеживаемые объекты, сохраняя и восстанавливая их.
tf.compat.v1.train.Checkpoint(
**kwargs
)
Конструктор Checkpoint принимает ключевые аргументы, значениями которых являются типы, содержащие отслеживаемое состояние, такие как реализации 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 (Layer, Optimizer, Variable и т. д.) с именованными ребрами, и этот граф используется для сопоставления переменных при восстановлении контрольной точки. Он может быть более устойчив к изменениям в программе Python и помогает поддерживать восстановление при создании переменных при выполнении eager-вычислений. Для нового кода предпочтительнее использовать tf.train.Checkpoint вместо tf.compat.v1.train.Saver.
Объекты 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.
Когда переменные присваиваются нескольким рабочим процессам, каждый рабочий процесс записывает свою секцию контрольной точки. Затем эти секции объединяются/переиндексируются, чтобы вести себя как одна контрольная точка. Это позволяет избежать копирования всех переменных в один рабочий процесс, но требует, чтобы все рабочие процессы видели одну общую файловую систему.
Хотя 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.
| Аргументы | |
|---|---|
**kwargs | Ключевые аргументы устанавливаются как атрибуты этого объекта и сохраняются вместе с контрольной точкой. Значения должны быть отслеживаемыми объектами. |
| Исключения | |
|---|---|
ValueError | Если объекты в kwargs не отслеживаются. |
| Атрибуты | |
|---|---|
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 как можно скорее.
| Аргументы | |
|---|---|
save_path | Путь к контрольной точке, возвращённый save или tf.train.latest_checkpoint. Если None (как при отсутствии последней контрольной точки для возврата tf.train.latest_checkpoint), возвращает объект, который может запускать инициализаторы для объектов в графе зависимостей. Если контрольная точка была записана с помощью tf.compat.v1.train.Saver на основе имён, для сопоставления переменных используются имена. |
| Возвращаемое значение | |
|---|---|
| Объект состояния загрузки, который может использоваться для утверждений о состоянии восстановления контрольной точки и выполнения операций инициализации/восстановления. Возвращаемый объект состояния имеет следующие методы:
|
save
save(
file_prefix, session=None, options=None
)
Сохраняет контрольную точку обучения и предоставляет базовый менеджмент контрольных точек.
Сохранённая контрольная точка включает переменные, созданные этим объектом, и любые отслеживаемые объекты, от которых он зависит на момент вызова Checkpoint.save().
save — это базовая обёртка над методом write, последовательно нумерующая контрольные точки с помощью save_counter и обновляющая метаданные, используемые tf.train.latest_checkpoint. Более продвинутый менеджмент контрольных точек, например, сборка мусора и настройка нумерации, может быть предоставлен другими утилитами, которые также обертывают write (tf.train.CheckpointManager, например).
| Аргументы | |
|---|---|
file_prefix | Префикс, используемый для имён файлов контрольных точек (/путь/к/каталогу/и_префикс). Имена генерируются на основе этого префикса и Checkpoint.save_counter. |
session | Сессия для оценки переменных. Игнорируется при выполнении с жадностью. Если не предоставлено при построении графа, используется сессия по умолчанию. |
options | Необязательный объект tf.train.CheckpointOptions. |
| Возвращаемое значение | |
|---|---|
| Полный путь к контрольной точке. |
write
write(
file_prefix, session=None, options=None
)
Записывает контрольную точку обучения.
Контрольная точка включает переменные, созданные этим объектом, и любые отслеживаемые объекты, от которых он зависит на момент вызова Checkpoint.write().
write не нумерует контрольные точки, не увеличивает save_counter и не обновляет метаданные, используемые tf.train.latest_checkpoint. Он предназначен в первую очередь для использования утилитами управления контрольными точками более высокого уровня. save предоставляет очень базовую реализацию этих функций.
| Аргументы | |
|---|---|
file_prefix | Префикс, используемый для имён файлов контрольных точек (/путь/к/каталогу/и_префикс). |
session | Сессия для оценки переменных. Игнорируется при выполнении с жадностью. Если не предоставлено при построении графа, используется сессия по умолчанию. |
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/compat/v1/train/Checkpoint