Spec-Zone.ru › TensorFlow 2.9

tf.experimental.dtensor.DTensorCheckpoint

Управляет сохранением/восстановлением значений, отслеживаемых в диске, для DTensor.

Наследуется от: Checkpoint

tf.experimental.dtensor.DTensorCheckpoint(
    mesh: tf.experimental.dtensor.Mesh,
    root=None,
    **kwargs
)
Аргументы
root Корневый объект для сохранения. root может быть отслеживаемым объектом или WeakRef отслеживаемого объекта.
**kwargs Аргументы ключевых слов устанавливаются как атрибуты этого объекта и сохраняются вместе с контрольной точкой. Все kwargs должны быть отслеживаемыми объектами или вложенной структурой отслеживаемых объектов (list, dict, или tuple).
Исключения
ValueError Если root или объекты в kwargs не отслеживаются. Исключение ValueError также генерируется, если объект root отслеживает разные объекты, чем те, которые перечислены в атрибутах в kwargs (например, root.child = A и tf.train.Checkpoint(root, child=B) несовместимы).
Атрибуты
save_counter Целочисленная переменная, которая начинается с нуля и увеличивается при сохранении.

Используется для нумерации контрольных точек.

Методы

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.train.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() либо сразу присваивает значения, если переменные для восстановления уже созданы, либо откладывает восстановление до создания переменных. Зависимости, добавленные после этого вызова, будут сопоставлены, если у них есть соответствующий объект в контрольной точке (запрос на восстановление будет поставлен в очередь в любом отслеживаемом объекте, ожидающем добавления ожидаемой зависимости).

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

# You can additionally pass options to restore():
options = tf.CheckpointOptions(experimental_io_device="/job:localhost")
checkpoint.restore(path, options=options)

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

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.train.Checkpoint(step=step)
checkpoint.save("/tmp/ckpt")

# Later, read the checkpoint with restore()
checkpoint.restore("/tmp/ckpt-1")

# You can also pass options to save() and restore(). For example this
# runs the IO ops on the localhost:
options = tf.train.CheckpointOptions(experimental_io_device="/job:localhost")
checkpoint.save("/tmp/ckpt", options=options)

# Later, read the checkpoint with restore()
checkpoint.restore("/tmp/ckpt-1", options=options)
Аргументы
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")

# 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)
Аргументы
file_prefix Префикс для имён файлов контрольных точек (/путь/к/каталогу/и_префикс).
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/versions/r2.9/api_docs/python/tf/experimental/dtensor/DTensorCheckpoint

Spec-Zone.ru

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