Функция recv
Функция recv из модуля torch.distributed выполняет
блокирующий приём тензора от указанного процесса-отправителя.
Она используется в распределённых вычислениях для обмена данными
между узлами. Первым параметром передаётся тензор, который будет
заполнен полученными данными. Вторым параметром указывается
источник - ранг процесса-отправителя. Третий параметр задаёт
группу процессов, по умолчанию используется группа по умолчанию.
Функция возвращает ранг процесса-отправителя.
Синтаксис
torch.distributed.recv(tensor, src, [group])
Пример с одним отправителем
Рассмотрим базовый пример приёма тензора от процесса с рангом 0:
import torch
import torch.distributed as dist
dist.init_process_group(backend='gloo')
rank = dist.get_rank()
if rank == 0:
t = torch.tensor([10, 20, 30, 40])
dist.send(t, dst=1)
elif rank == 1:
res = torch.zeros(4, dtype=torch.long)
src_rank = dist.recv(res, src=0)
print(f"Received from {src_rank}: {res}")
dist.destroy_process_group()
Результат выполнения кода для процесса с рангом 1:
Received from 0: tensor([10, 20, 30, 40])
Пример с динамическим источником
Если передать src равным dist.ANY_SOURCE,
то функция примет тензор от любого процесса. Возвращаемое
значение покажет, от кого именно пришли данные:
import torch
import torch.distributed as dist
dist.init_process_group(backend='gloo')
rank = dist.get_rank()
size = dist.get_world_size()
if rank == 0:
t = torch.tensor([1, 2, 3])
dist.send(t, dst=2)
elif rank == 1:
t = torch.tensor([4, 5, 6])
dist.send(t, dst=2)
elif rank == 2:
res = torch.zeros(3, dtype=torch.long)
src_rank = dist.recv(res, src=dist.ANY_SOURCE)
print(f"Received from {src_rank}: {res}")
dist.destroy_process_group()
Результат выполнения кода для процесса с рангом 2:
Received from 0: tensor([1, 2, 3])
Пример с указанием группы
Функция позволяет работать внутри пользовательской группы процессов. Создадим подгруппу и выполним приём внутри неё:
import torch
import torch.distributed as dist
dist.init_process_group(backend='gloo')
rank = dist.get_rank()
subgroup = dist.new_group(ranks=[0, 1])
if rank == 0:
t = torch.tensor([7, 8, 9])
dist.send(t, dst=1, group=subgroup)
elif rank == 1:
res = torch.zeros(3)
src_rank = dist.recv(res, src=0, group=subgroup)
print(f"From {src_rank} in subgroup: {res}")
dist.destroy_process_group()
Результат выполнения кода для процесса с рангом 1:
From 0 in subgroup: tensor([7., 8., 9.])