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, qint16, quint16, uint16, complex128, half, uint32, uint64, bool. Локальный вход для суммы. |
group_assignment | A Tensor типа int32. Массив int32 с формой [num_groups, num_replicas_per_group]. group_assignment[i] представляет идентификаторы реплик в i-ой подгруппе. |
concat_dimension | A int. Номер измерения для конкатенации. |
split_dimension | A int. Номер измерения для разделения. |
split_count | A int. Количество разделов; это число должно быть равно размеру подгруппы (group_assignment.get_shape()[1]) |
name | Имя операции (необязательно). |
| Возвращаемое значение | |
|---|---|
A Tensor. Имеет тот же тип, что и input. |
© 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/raw_ops/AllToAll