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