tf.data.experimental.CheckpointInputPipelineHook
| Просмотреть исходный код на GitHub |
Сохраняет состояние конвейера входных данных каждые N шагов или секунд.
Наследуется от: SessionRunHook
tf.data.experimental.CheckpointInputPipelineHook(
estimator, external_state_policy='fail'
)
Этот хук сохраняет состояние итераторов в 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, которая скоро будет закрыта. |
© 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/r2.3/api_docs/python/tf/data/experimental/CheckpointInputPipelineHook