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(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, либо сразу присваивает значения, если переменные для восстановления уже созданы, либо откладывает восстановление до создания переменных. Зависимости, добавленные после этого вызова, будут сопоставлены, если в контрольной точке есть соответствующий объект (запрос на восстановление будет помещен в очередь в любой отслеживаемый объект, ожидающий добавления ожидаемой зависимости).
При построении графа операции восстановления добавляются в граф, но не выполняются сразу.
Чтобы убедиться, что загрузка завершена и больше не будет происходить назначений, используйте метод assert_consumed() объекта состояния, возвращенного restore:
checkpoint = tf.train.Checkpoint( ... ) checkpoint.restore(path).assert_consumed()
Будет возбуждено исключение, если какие-либо объекты Python в графе зависимостей не были найдены в контрольной точке или если какие-либо сохранённые значения не имеют соответствующего объекта Python.
При построении графа assert_consumed() указывает, что все операции восстановления, которые будут созданы для этой контрольной точки, созданы. Они могут быть выполнены с помощью метода run_restore_ops() объекта состояния:
checkpoint.restore(path).assert_consumed().run_restore_ops()
Если контрольная точка не была полностью обработана, список операций восстановления будет увеличиваться по мере добавления большего количества объектов в граф зависимостей.
Контрольные точки 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 на основе имён, используются имена для сопоставления переменных. |
| Returns | |
|---|---|
| Объект статуса загрузки, который может использоваться для утверждений о состоянии восстановления контрольной точки и выполнения операций инициализации/восстановления. Возвращаемый объект статуса имеет следующие методы:
|
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 | Сессия для вычисления переменных. Игнорируется при выполнении жадно. Если не предоставлена при построении графа, используется сессия по умолчанию. |
| Возвращаемое значение | |
|---|---|
| Полный путь к контрольной точке. |
write
write(
file_prefix, session=None
)
Записывает контрольную точку обучения.
Контрольная точка включает переменные, созданные этим объектом, и любые отслеживаемые объекты, от которых он зависит на момент вызова Checkpoint.write().
write не нумерует контрольные точки, не увеличивает save_counter, и не обновляет метаданные, используемые tf.train.latest_checkpoint. Она предназначена в первую очередь для использования средствами управления контрольными точками более высокого уровня. save предоставляет очень базовую реализацию этих функций.
| Аргументы | |
|---|---|
file_prefix | Префикс, используемый для имён файлов контрольных точек (/путь/к/каталогу/и_префикс). |
session | Сессия для вычисления переменных. Игнорируется при выполнении жадно. Если не предоставлена при построении графа, используется сессия по умолчанию. |
| Возвращаемое значение | |
|---|---|
Полный путь к контрольной точке (т.е. 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/compat/v1/train/Checkpoint