Функция reduce_scatter_tensor
Функция distributed.reduce_scatter_tensor выполняет комбинированную операцию: сначала выполняет редукцию (например, суммирование) переданных от всех процессов тензоров, а затем разбрасывает результирующие фрагменты между процессами. Эта функция полезна для эффективного обмена данными в распределённых системах, где требуется агрегировать информацию и распределить её обратно. В отличие от reduce_scatter, эта функция работает с тензорами как с целостными объектами, а не со списками.
Первый параметр output - тензор для приёма части редуцированного результата. Второй параметр input - тензор, содержащий данные для редукции. Третий параметр op - операция редукции (по умолчанию ReduceOp.SUM). Четвёртый параметр group - группа процессов (по умолчанию группа по умолчанию).
Синтаксис
torch.distributed.reduce_scatter_tensor(
output,
input,
op=ReduceOp.SUM,
group=None
)
Пример
Давайте выполним редукцию с разбросом в группе из двух процессов:
import torch
import torch.distributed as dist
torch.manual_seed(0)
dist.init_process_group("gloo")
rank = dist.get_rank()
world_size = dist.get_world_size()
input_tensor = torch.tensor([rank + 1, rank + 2])
output_tensor = torch.zeros(world_size)
dist.reduce_scatter_tensor(
output_tensor,
input_tensor,
op=dist.ReduceOp.SUM
)
if rank == 0:
print(f"Rank {rank} output: {output_tensor}")
dist.barrier()
if rank == 1:
print(f"Rank {rank} output: {output_tensor}")
dist.destroy_process_group()
Результат выполнения кода для процесса 0:
Rank 0 output: tensor([2., 4.])
Результат выполнения кода для процесса 1:
Rank 1 output: tensor([0., 0.])
Пример
Выполним редукцию с операцией PRODUCT:
import torch
import torch.distributed as dist
torch.manual_seed(0)
dist.init_process_group("gloo")
rank = dist.get_rank()
world_size = dist.get_world_size()
input_tensor = torch.tensor([rank + 1.0, rank + 2.0])
output_tensor = torch.zeros(world_size)
dist.reduce_scatter_tensor(
output_tensor,
input_tensor,
op=dist.ReduceOp.PRODUCT
)
if rank == 0:
print(f"Rank {rank} output: {output_tensor}")
dist.destroy_process_group()
Результат выполнения кода для процесса 0:
Rank 0 output: tensor([2., 6.])
Смотрите также
-
функцию
reduce_scatter,
которая выполняет редукцию и разброс списка тензоров -
функцию
all_reduce,
которая выполняет редукцию тензора на всех процессах -
функцию
all_gather,
которая собирает тензоры со всех процессов на каждый процесс -
перечисление
ReduceOp,
которое определяет доступные операции редукции