Spec-Zone.ru › TensorFlow 2.9

tf.data.experimental.CheckpointInputPipelineHook

Просмотреть исходный код на GitHub

Сохраняет состояние входной цепочки обработки данных каждые N шагов или секунд.

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

Просмотр псевдонимов

Псевдонимы для миграции

См. Руководство по миграции для получения дополнительных сведений.

tf.compat.v1.data.experimental.CheckpointInputPipelineHook

tf.data.experimental.CheckpointInputPipelineHook(
    estimator, external_state_policy=None
)

Этот хук сохраняет состояние итераторов в Graph, чтобы при возобновлении обучения входная цепочка обработки данных продолжила работу с того места, где она остановилась. Это может потенциально предотвратить переобучение в некоторых цепочках, где количество шагов обучения на один вызов оценки невелико по сравнению с размером набора данных, или если цепочка обучения прервана.

Отличия от CheckpointSaverHook:

  1. Сохраняет только входные цепочки обработки данных в коллекции "итераторы", а не глобальные переменные или другие сохраняемые объекты.
  2. Не записывает GraphDef и MetaGraphDef в журнал.

Пример сохранения точки контроля цепочки обучения:

est = tf.estimator.Estimator(model_fn)
while True:
  est.train(
      train_input_fn,
      hooks=[tf.data.experimental.CheckpointInputPipelineHook(est)],
      steps=train_steps_per_eval)
  # Note: We do not pass the hook here.
  metrics = est.evaluate(eval_input_fn)
  if should_stop_the_training(metrics):
    break

Этот хук следует использовать, если состояние входной цепочки обработки данных необходимо сохранить отдельно от точки контроля модели. Это может быть полезно по нескольким причинам:

  1. Точка контроля входной цепочки обработки данных может быть большой, если есть большие буферы перемешивания или предварительной выборки, и может увеличить размер точки контроля.
  2. Если входная цепочка обработки данных используется совместно для обучения и проверки, восстановление точки контроля во время проверки может перезаписать входную цепочку обработки данных для проверки.

Для сохранения точки контроля входной цепочки обработки данных вместе с весами модели используйте tf.data.experimental.make_saveable_from_iterator напрямую для создания SaveableObject и добавления в коллекцию SAVEABLE_OBJECTS. Однако имейте в виду, что вам необходимо будет следить за тем, чтобы во время проверки не восстанавливалась итератор обучения. Вы можете сделать это, не добавляя итератор в коллектор SAVEABLE_OBJECTS при построении графа оценки.

Аргументы
estimator Обучающий прибор.
external_state_policy Строка, определяющая, как обрабатывать входные цепочки обработки данных, зависящие от внешнего состояния. Возможные значения: 'ignore': Внешнее состояние игнорируется. 'warn': Внешнее состояние игнорируется с выводом предупреждения. 'fail': Операция завершается с ошибкой при обнаружении внешнего состояния. По умолчанию установлено 'fail'.
Исключения
ValueError Должно быть установлено одно из save_steps или save_secs.
ValueError Должно быть установлено не более одного из saver или scaffold.
ValueError Если external_state_policy не равно 'warn', 'ignore' или 'fail'.

Методы

after_create_session

Просмотреть исходный код

after_create_session(
    session, coord
)

Вызывается при создании новой сессии TensorFlow.

Это вызывается для уведомления хуков о создании новой сессии. Есть два основных отличия от ситуации, когда вызывается begin:

  • При вызове этого метода граф завершается, и больше невозможно добавлять операции в граф.
  • Этот метод также вызывается в результате восстановления обернутой сессии, а не только в начале всей сессии.
Аргументы
session Созданная сессия TensorFlow.
coord Объект Coordinator, который отслеживает все потоки.

after_run

Просмотреть исходный код

after_run(
    run_context, run_values
)

Вызывается после каждого вызова run().

Аргумент run_values содержит результаты запрошенных операций/тензоров функцией before_run().

Аргумент run_context такой же, как и переданный в вызов before_run. Вызов run_context.request_stop() может быть использован для остановки итерации.

Если вызов session.run() вызывает исключения, то after_run() не вызывается.

Аргументы
run_context Объект SessionRunContext.
run_values Объект SessionRunValues.

before_run

Просмотреть исходный код

before_run(
    run_context
)

Вызывается перед каждым вызовом run().

Вы можете вернуть из этого вызова объект SessionRunArgs, указывающий операции или тензоры, которые нужно добавить к предстоящему вызову run(). Эти операции/тензоры будут выполнены вместе с операциями/тензорами, изначально переданными в исходный вызов run(). Аргументы вызова run, которые вы возвращаете, также могут содержать данные для добавления в вызов run().

Аргумент run_context — это объект SessionRunContext, который предоставляет информацию о предстоящем вызове run(): первоначально запрошенные операции/тензоры, сессия TensorFlow.

На этом этапе граф завершён и вы не можете добавлять операции.

Аргументы
run_context Объект SessionRunContext.
Возвращаемое значение
None или объект SessionRunArgs.

begin

Просмотреть исходный код

begin()

Вызывается один раз перед использованием сессии.

При вызове текущим графом является тот, который будет запущен в сессии. Хук может изменять граф, добавляя новые операции. После вызова begin() граф будет завершен, и другие обратные вызовы больше не смогут изменять граф. Второй вызов begin() на том же графе не должен изменять граф.

end

Просмотреть исходный код

end(
    session
)

Вызывается в конце сессии.

Аргумент session может использоваться в случае, если хук хочет выполнить окончательные операции, такие как сохранение последней точки контроля.

Если session.run() вызывает исключение, отличное от OutOfRangeError или StopIteration, то end() не вызывается. Обратите внимание на разницу между поведением end() и after_run() при вызове session.run() исключения OutOfRangeError или StopIteration. В этом случае end() вызывается, но after_run() не вызывается.

Аргументы
session Сессия TensorFlow, которая скоро будет закрыта.

© 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/data/experimental/CheckpointInputPipelineHook

Spec-Zone.ru

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