tf.data.experimental.CheckpointInputPipelineHook
| Просмотреть исходный код на GitHub |
Сохраняет состояние входной цепочки обработки данных каждые N шагов или секунд.
Наследуется от: SessionRunHook
tf.data.experimental.CheckpointInputPipelineHook(
estimator, external_state_policy=None
)
Этот хук сохраняет состояние итераторов в Graph, чтобы при возобновлении обучения входная цепочка обработки данных продолжила работу с того места, где она остановилась. Это может потенциально предотвратить переобучение в некоторых цепочках, где количество шагов обучения на один вызов оценки невелико по сравнению с размером набора данных, или если цепочка обучения прервана.
Отличия от CheckpointSaverHook:
- Сохраняет только входные цепочки обработки данных в коллекции "итераторы", а не глобальные переменные или другие сохраняемые объекты.
- Не записывает
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
Этот хук следует использовать, если состояние входной цепочки обработки данных необходимо сохранить отдельно от точки контроля модели. Это может быть полезно по нескольким причинам:
- Точка контроля входной цепочки обработки данных может быть большой, если есть большие буферы перемешивания или предварительной выборки, и может увеличить размер точки контроля.
- Если входная цепочка обработки данных используется совместно для обучения и проверки, восстановление точки контроля во время проверки может перезаписать входную цепочку обработки данных для проверки.
Для сохранения точки контроля входной цепочки обработки данных вместе с весами модели используйте 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