Spec-Zone.ru › TensorFlow 1.15

tf.data.experimental.CheckpointInputPipelineHook

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

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

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

Псевдонимы

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

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

tf.compat.v1.data.experimental.CheckpointInputPipelineHook, `tf.compat.v2.data.experimental.CheckpointInputPipelineHook`

tf.data.experimental.CheckpointInputPipelineHook(
    estimator
)

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

Отличия от CheckpointSaverHook:

  1. Сохраняет только конвейеры ввода данных в коллекции "iterators", а не глобальные переменные или другие сохраняемые объекты.
  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 при построении графика оценки.

Args
estimator Estimator.
Raises
ValueError Один из save_steps или save_secs должен быть задан.
ValueError Не более одного из saver или scaffold должен быть задан.

Методы

after_create_session

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

after_create_session(
    session, coord
)

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

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

  • Когда это вызывается, граф финализирован, и операции больше нельзя добавлять в граф.
  • Этот метод также будет вызван в результате восстановления обернутой сессии, а не только в начале всей сессии.
Args
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() не вызывается.

Args
run_context Объект SessionRunContext
run_values Объект SessionRunValues.

before_run

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

before_run(
    run_context
)

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

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

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

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

Args
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() не вызывается.

Args
session Закрываемая TensorFlow сессия.

© 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/r1.15/api_docs/python/tf/data/experimental/CheckpointInputPipelineHook

Spec-Zone.ru

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