tf.raw_ops.AllToAll
Операция для обмена данными между репликами TPU.
tf.raw_ops.AllToAll(
input, group_assignment, concat_dimension, split_dimension, split_count,
name=None
)
В каждой реплике входные данные разбиваются на split_count блока по split_dimension и отправляются другим репликам с учётом group_assignment. После получения split_count - 1 блоков от других реплик, блоки конкатенируются по concat_dimension в качестве выходных данных.
Например, предположим, что есть 2 реплики TPU: реплика 0 получает вход: [[A, B]] реплика 1 получает вход: [[C, D]]
group_assignment=[[0, 1]] concat_dimension=0 split_dimension=1 split_count=2
выход реплики 0: [[A], [C]] выход реплики 1: [[B], [D]]
| Аргументы | |
|---|---|
input | A Tensor. Должен быть одним из следующих типов: float32, float64, int32, uint8, int16, int8, complex64, int64, qint8, quint8, qint32, bfloat16, uint16, complex128, half, uint32, uint64, bool. Локальный вход в сумму. |
group_assignment | A Tensor типа int32. Tensor int32 с формой [num_groups, num_replicas_per_group]. group_assignment[i] представляет идентификаторы реплик в i-ой подгруппе. |
concat_dimension | An int. Номер измерения для конкатенации. |
split_dimension | An int. Номер измерения для разделения. |
split_count | An int. Количество разделений, это число должно быть равно размеру подгруппы (group_assignment.get_shape()[1]) |
name | Имя операции (необязательно). |
| Возвращаемые значения | |
|---|---|
A Tensor. Имеет тот же тип, что и input. |
© 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.4/api_docs/python/tf/raw_ops/AllToAll