tf.distribute.experimental.PreemptionCheckpointHandler
Обработчик прерывания и ошибок для синхронной тренировки.
tf.distribute.experimental.PreemptionCheckpointHandler(
cluster_resolver,
checkpoint_or_checkpoint_manager,
checkpoint_dir=None,
termination_config=None
)
Примечание: Этот API поддерживает только использование сtf.distribute.MultiWorkerMirroredStrategyиtf.distribute.TPUStrategy.
Объект PreemptionCheckpointHandler координирует все рабочие узлы для сохранения контрольной точки при получении сигнала о прерывании. Он также помогает точно распространять сообщения об ошибках приложения по кластеру. При создании объекта PreemptionCheckpointHandler он восстанавливает значения из последнего файла контрольной точки, если он существует.
Сразу после инициализации объект начинает следить за сигналами завершения для любого узла в кластере. При получении сигнала, в следующий раз, когда узел выполняет PreemptionCheckpointHandler.run, объект PreemptionCheckpointHandler синхронизирует все узлы для сохранения контрольной точки. Затем, если настроен обработчик exit_fn через tf.distribute.experimental.TerminationConfig, он вызывается. В противном случае процесс просто завершается, а платформа должна его перезапустить позже.
Примечание: Мы рекомендуем пользователямtf.distribute.MultiWorkerMirroredStrategy, которые настраивают собственный обработчикexit_fnвtf.distribute.experimental.TerminationConfig, включать в негоsys.exit(CODE_OR_MESSAGE)вexit_fn, чтобы после перезапуска все узлы могли правильно инициализировать службы связи. Для пользователейtf.distribute.TPUStrategy, если они не хотят перезапуска кластера, но хотят перезапуск внутри процесса (то есть сохранить координатор активным и повторить шаги подключения к кластеру, инициализации системы TPU и создания объектаTPUStrategy), они могут настроить обработчикexit_fnкак бездействие.
Для пользователей tf.distribute.MultiWorkerMirroredStrategy основной API является PreemptionCheckpointHandler.run:
strategy = tf.distribute.MultiWorkerMirroredStrategy()
trained_epoch = tf.Variable(initial_value=tf.constant(0, dtype=tf.dtypes.int64), name='epoch')
step_in_epoch = tf.Variable(initial_value=tf.constant(0, dtype=tf.dtypes.int64), name='step_in_epoch')
with strategy.scope():
dataset, model, optimizer = ...
checkpoint = tf.train.Checkpoint(optimizer=optimizer,
model=model,
trained_epoch=trained_epoch,
step_in_epoch=step_in_epoch)
preemption_checkpoint_handler = tf.distribute.experimental.PreemptionCheckpointHandler(cluster_resolver, checkpoint, checkpoint_dir)
while trained_epoch.numpy() < NUM_EPOCH:
while step_in_epoch.numpy() < STEPS_PER_EPOCH:
# distributed_train_function contains a call to strategy.run.
loss += preemption_checkpoint_handler.run(distributed_train_function, args=(next(iterator),))
# For users of MultiWorkerMirroredStrategy, usually
# STEPS_PER_TRAIN_FUNCTION = 1.
step_in_epoch.assign_add(STEPS_PER_TRAIN_FUNCTION)
...
epoch.assign_add(1)
step_in_epoch.assign(0)
Для пользователей tf.distribute.TPUStrategy основными API являются PreemptionCheckpointHandler.run и PreemptionCheckpointHandler.watch_preemption_scope:
strategy = tf.distribute.TPUStrategy(tpu_cluster_resolver)
# Rest of TPU init omitted, see documentation for TPUSTrategy.
with preemption_checkpoint_handler.watch_preemption_scope():
while trained_epoch.numpy() < NUM_EPOCH:
while step_in_epoch.numpy() < STEPS_PER_EPOCH:
# distributed_train_function contains a call to strategy.run.
loss += preemption_checkpoint_handler.run(distributed_train_function, args=(next(iterator),))
# For users of TPUStrategy, usually STEPS_PER_TRAIN_FUNCTION >> 1 since
# clustering multiple steps within a tf.function amortizes the overhead
# of launching a multi-device function on TPU Pod.
step_in_epoch.assign_add(STEPS_PER_TRAIN_FUNCTION)
...
epoch.assign_add(1)
step_in_epoch.assign(0)
Не все прерывания сопровождаются предварительным уведомлением, чтобы PreemptionCheckpointHandler мог их обработать, например, те, которые вызваны отказом оборудования. Для пользователя, который самостоятельно сохраняет контрольные точки для этих случаев вне PreemptionCheckpointHandler, если он использует tf.train.CheckpointManager, передайте его в качестве аргумента checkpoint_or_checkpoint_manager методу PreemptionCheckpointHandler. Если у него нет tf.train.CheckpointManager, но он непосредственно работает с tf.train.Checkpoint, мы рекомендуем сохранять контрольные точки в каталоге, переданном в качестве аргумента checkpoint_dir. Таким образом, в начале программы PreemptionCheckpointHandler может восстановить последнюю контрольную точку из каталога, независимо от того, сохранил ли её сам пользователь или PreemptionCheckpointHandler до прерывания.
Примечание о платформе:
PreemptionCheckpointHandler может обрабатывать только виды прерывания с предварительным уведомлением. На данный момент API распознаёт сигналы завершения для CPU, GPU и TPU на Google Borg и CPU и GPU на Google Cloud Platform. В этих случаях PreemptionCheckpointHandler автоматически использует соответствующий механизм обнаружения уведомлений о прерывании/техническом обслуживании. Пользователи других платформ могут настроить поведение мониторинга обнаружения через tf.distribute.experimental.TerminationConfig. Здесь также можно настроить поведение выхода и длительность периода ожидания.
| Аргументы | |
|---|---|
cluster_resolver | объект tf.distribute.cluster_resolver.ClusterResolver. Вы также можете получить его через атрибут cluster_resolver используемой стратегии распределения. |
checkpoint_or_checkpoint_manager | tf.train.CheckpointManager или tf.train.Checkpoint. Если вы используете tf.train.CheckpointManager для управления контрольными точками вне PreemptionCheckpointHandler в целях резервного копирования, передайте его как аргумент checkpoint_or_checkpoint_manager. В противном случае передайте tf.train.Checkpoint и PreemptionCheckpointHandler создаст tf.train.CheckpointManager для его управления в checkpoint_dir. |
checkpoint_dir | каталог, в котором PreemptionCheckpointHandler сохраняет и восстанавливает контрольные точки. При создании PreemptionCheckpointHandler восстанавливается последняя контрольная точка в checkpoint_dir. (Это не требуется, если вместо tf.train.Checkpoint передаётся tf.train.CheckpointManager в качестве аргумента checkpoint_or_checkpoint_manager.) |
termination_config | необязательно, объект tf.distribute.experimental.TerminationConfig для настройки платформы, отличной от Google Borg или GCP. |
Методы
run
run(
distributed_train_function, *args, **kwargs
)
Выполняет функцию обучения с обработкой ошибок и прерываний.
Эта функция обрабатывает сигнал прерывания от любого узла в кластере, сохраняя прогресс обучения и завершая работу корректно. Она также распространит любую ошибку программы, возникшую во время выполнения distributed_train_function, на все рабочие узлы, чтобы они могли вызвать ту же ошибку.
Аргумент distributed_train_function должен быть функцией распределённого обучения (т.е. содержащей вызов tf.distribute.Strategy.run). Пользователям tf.distribute.MultiWorkerMirroredStrategy рекомендуется передавать в PreemptionCheckpointHandler.run одношаговую функцию distributed_train_function, чтобы контрольная точка могла быть сохранена вовремя в случае отправки сигнала прерывания или уведомления о техническом обслуживании.
Помимо обработки прерываний и ошибок, PreemptionCheckpointHandler.run(distributed_train_function, *args, **kwargs) имеет тот же эффект и выход, что и distributed_train_function(*args, **kwargs). distributed_train_function может возвращать некоторые или никакие результаты. Ниже приведён сокращённый пример:
@tf.function
def distributed_train_step(iterator):
# A distributed single-step training function.
def step_fn(inputs):
# A per-replica single-step training function.
x, y = inputs
...
return loss
per_replica_losses = strategy.run(step_fn, args=(next(iterator),))
return strategy.reduce(
tf.distribute.ReduceOp.SUM, per_replica_losses, axis=None)
for epoch in range(preemption_handler.total_run_calls // STEPS_PER_EPOCH,
EPOCHS_TO_RUN):
iterator = iter(multi_worker_dataset)
total_loss = 0.0
num_batches = 0
for step in range(preemption_handler.total_run_calls % STEPS_PER_EPOCH,
STEPS_PER_EPOCH):
total_loss += preemption_handler.run(distributed_train_step)
num_batches += 1
train_loss = total_loss / num_batches
print('Epoch: %d, train_loss: %f.' %(epoch.numpy(), train_loss))
train_accuracy.reset_states()
| Аргументы | |
|---|---|
distributed_train_function | Функция распределённого обучения (одношаговая). |
*args | аргументы для distributed_train_function. |
**kwargs | ключевые слова аргументов для distributed_train_function. |
| Исключения | |
|---|---|
Ошибка программы, возникшая на любом узле кластера во время выполнения distributed_train_function, или любая ошибка при распространении ошибки. |
| Возвращаемое значение | |
|---|---|
Результат выполнения distributed_train_function. |
save_checkpoint_if_preempted
save_checkpoint_if_preempted(
*args, **kwargs
)
Сохраняет контрольную точку, если сигнал прерывания был получен.
Это альтернативный API для PreemptionCheckpointHandler.run и PreemptionCheckpointHandler.watch_preemption_scope. Этот метод работает как для tf.distribute.MultiWorkerMirroredStrategy, так и для tf.distribute.TPUStrategy. Однако, для TPUStrategy этот метод добавит синхронную точку между узлами и координатором, что может повлиять на производительность. Если это вызывает озабоченность, используйте сочетание PreemptionCheckpointHandler.watch_preemption_scope и PreemptionCheckpointHandler.run.
strategy = tf.distribute.TPUStrategy(tpu_cluster_resolver) # initialization omitted with strategy.scope(): # Save in the checkpoint. trained_step = tf.Variable(initial_value=tf.constant(0, dtype=tf.dtypes.int64), name='trained_step', aggregation=tf.VariableAggregation.ONLY_FIRST_REPLICA) checkpoint_manager = tf.train.CheckpointManager(checkpoint, directory, max_to_keep=1) preemption_handler = tf.distribute.experimental.PreemptionCheckpointHandler(cluster_resolver, checkpoint_manager) while trained_step.numpy() < NUM_STEPS: # Train STEPS_IN_FUNCTION steps at once. train_multi_step_function() trained_step.assign_add(STEPS_IN_FUNCTION) preemption_handler.save_checkpoint_if_preempted()
| Аргументы | |
|---|---|
*args | аргументы для tf.train.CheckpointManager.save() для сохранения контрольной точки. |
**kwargs | ключевые слова аргументов для tf.train.CheckpointManager.save() для сохранения. |
watch_preemption_scope
@tf_contextlib.contextmanager watch_preemption_scope()
Синхронизирует ошибки и, возможно, сохраняет контрольную точку для использования со стратегией TPUStrategy.
Примечание: Для использования с tf.distribute.MultiWorkerMirroredStrategy этот API не требуется.
Пример использования:
with preemption_checkpoint_handler.watch_preemption_scope():
while trained_step.numpy() < NUM_STEPS:
# distributed_train_function contains a call to strategy.run.
loss += preemption_checkpoint_handler.run(distributed_train_function, args=(next(iterator),))
trained_step.assign_add(STEPS_PER_TRAIN_FUNCTION)
В этом потоке работы PreemptionCheckpointHandler.run пометит полученный сигнал прерывания, а watch_preemption_scope обработает сигнал прерывания, сохранив контрольную точку, а затем выйдет для перезапуска или выполнит переданную пользователем exit_fn в tf.distribute.experimental.TerminationConfig. Если сигнал прерывания не получен во время выполнения операций и функций в рамках области, watch_preemption_scope гарантирует завершение всех асинхронных операций и функций при выходе и вызовет исключения, если асинхронное выполнение приведет к ошибочному состоянию.
| Возвращаемое значение | |
|---|---|
| Нет |
© 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/api_docs/python/tf/distribute/experimental/PreemptionCheckpointHandler