Функция all_gather
Функция torch.distributed.all_gather является ключевым инструментом в распределенных вычислениях PyTorch.
Она выполняет коллективную операцию, при которой каждый процесс отправляет свой тензор всем остальным процессам в группе.
В результате каждый процесс получает список тензоров, содержащий данные от всех участников группы.
Порядок тензоров в списке соответствует порядку рангов процессов в группе.
Это позволяет синхронизировать состояние модели или обмениваться промежуточными результатами вычислений.
Функция требует, чтобы все тензоры были одинакового размера и типа. Первым параметром передается объект группы (группа процессов), вторым - список тензоров для записи результатов, третьим - тензор с данными от текущего процесса.
Синтаксис
torch.distributed.all_gather(
group: ProcessGroup,
tensor_list: List[Tensor],
tensor: Tensor
) -> None
Пример
Перед использованием распределенных функций необходимо инициализировать группу процессов. Давайте создадим простой пример, где каждый процесс отправляет свой ранг всем остальным:
import torch
import torch.distributed as dist
# Инициализация распределенной среды (обычно выполняется при запуске)
dist.init_process_group(backend='gloo')
rank = dist.get_rank()
world_size = dist.get_world_size()
# Тензор для отправки (содержит ранг процесса)
t = torch.tensor([rank], dtype=torch.int64)
# Список для приема данных от всех процессов
tensor_list = [torch.zeros(1, dtype=torch.int64) for _ in range(world_size)]
# Сбор данных со всех процессов
dist.all_gather(tensor_list, t)
if rank == 0:
print(tensor_list)
dist.destroy_process_group()
Результат выполнения кода на процессе с рангом 0 в группе из 4 процессов:
[tensor([0]), tensor([1]), tensor([2]), tensor([3])]
Пример
Следующий пример демонстрирует сбор многомерных тензоров. Каждый процесс создает случайный тензор размерности 2x3 и отправляет его всем остальным:
import torch
import torch.distributed as dist
torch.manual_seed(42)
dist.init_process_group(backend='gloo')
rank = dist.get_rank()
world_size = dist.get_world_size()
# Каждый процесс создает свой тензор
t = torch.randn(2, 3) + rank
# Подготовка списка для приема
tensor_list = [torch.zeros_like(t) for _ in range(world_size)]
# Сбор всех тензоров
dist.all_gather(tensor_list, t)
if rank == 0:
for i, tensor in enumerate(tensor_list):
print(f"Тензор от процесса {i}:")
print(tensor)
dist.destroy_process_group()
Пример
При работе с NCCL бэкендом тензоры должны находиться на GPU. Давайте рассмотрим пример сбора данных на GPU:
import torch
import torch.distributed as dist
dist.init_process_group(backend='nccl')
rank = dist.get_rank()
world_size = dist.get_world_size()
torch.cuda.set_device(rank)
# Тензор на GPU
t = torch.tensor([rank * 10], dtype=torch.float32, device='cuda')
# Список тензоров на GPU
tensor_list = [
torch.zeros(1, dtype=torch.float32, device='cuda')
for _ in range(world_size)
]
dist.all_gather(tensor_list, t)
if rank == 0:
res = [tensor.cpu() for tensor in tensor_list]
print(res)
dist.destroy_process_group()
Результат выполнения кода на процессе с рангом 0 в группе из 4 процессов:
[tensor([0.]), tensor([10.]), tensor([20.]), tensor([30.])]
Смотрите также
-
функцию
all_reduce,
которая выполняет редукцию данных от всех процессов -
функцию
broadcast,
которая отправляет тензор от одного процесса всем остальным -
функцию
reduce_scatter,
которая выполняет редукцию и рассылает результаты по частям -
функцию
all_gather_into_tensor,
которая собирает данные в один объединенный тензор