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