tf.distribute.ReplicaContext
| Просмотр исходного кода на GitHub |
tf.distribute.Strategy API, когда вы находитесь в контексте реплики.
tf.distribute.ReplicaContext(
strategy, replica_id_in_sync_group
)
Вы можете использовать tf.distribute.get_replica_context, чтобы получить экземпляр ReplicaContext. Это должно находиться внутри вашей функции шага репликации, например, в вызове tf.distribute.Strategy.run.
| Атрибуты | |
|---|---|
devices | Устройства, на которых будет выполняться эта реплика, в виде кортежа строк. |
num_replicas_in_sync | Возвращает число реплик, по которым агрегируются градиенты. |
replica_id_in_sync_group | Возвращает идентификатор реплики, которая определяется. Это идентифицирует реплику, которая является частью группы синхронизации. В настоящее время мы предполагаем, что все группы синхронизации содержат одинаковое количество реплик. Значение идентификатора реплики может варьироваться от 0 до Примечание: Это не гарантирует, что это тот же идентификатор, что и идентификатор реплики XLA, используемый для низкоуровневых операций, таких как collective_permute. |
strategy | Текущий объект tf.distribute.Strategy. |
Методы
all_reduce
all_reduce(
reduce_op, value, experimental_hints=None
)
Выполняет всеохватывающее сокращение заданного value Tensor вложенного значения по всем репликам.
Если all_reduce вызывается в любой реплике, то его необходимо вызвать во всех репликах. Вложенная структура и Tensor формы должны быть идентичными во всех репликах.
Пример с двумя репликами: Реплика 0 value: {'a': 1, 'b': [40, 1]} Реплика 1 value: {'a': 3, 'b': [ 2, 98]}
Если reduce_op == SUM: Результат (на всех репликах): {'a': 4, 'b': [42, 99]}
Если reduce_op == MEAN: Результат (на всех репликах): {'a': 2, 'b': [21, 49.5]}
| Аргументы | |
|---|---|
reduce_op | Тип сокращения, экземпляр перечисления tf.distribute.ReduceOp. |
value | Вложенная структура Tensor для всеохватывающего сокращения. Структура должна быть совместима с tf.nest. |
experimental_hints | tf.distrbute.experimental.CollectiveHints. Указания для выполнения коллективных операций. |
| Возвращает | |
|---|---|
Вложенное Tensor со сжатыми value с каждой реплики. |
merge_call
merge_call(
merge_fn, args=(), kwargs=None
)
Объединяет аргументы по репликам и выполняет merge_fn в контексте между репликами.
Это позволяет обмениваться данными и координировать действия, когда возникает несколько вызовов step_fn, которые вызваны вызовом strategy.run(step_fn, ...).
См. tf.distribute.Strategy.run для объяснения.
Если не внутри распределённого контекста, это эквивалентно:
strategy = tf.distribute.get_strategy() with cross-replica-context(strategy): return merge_fn(strategy, *args, **kwargs)
| Аргументы | |
|---|---|
merge_fn | Функция, объединяющая аргументы из потоков, которые передаются как PerReplica. Она принимает объект tf.distribute.Strategy в качестве первого аргумента. |
args | Список или кортеж с позиционными аргументами для каждого потока для merge_fn. |
kwargs | Словарь с ключевыми аргументами для каждого потока для merge_fn. |
| Возвращает | |
|---|---|
Значение возврата merge_fn, за исключением значений PerReplica которые распаковываются. |
__enter__
__enter__()
__exit__
__exit__(
exception_type, exception_value, traceback
)
© 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/distribute/ReplicaContext