Функция distributed.broadcast
Функция distributed.broadcast выполняет широковещательную рассылку тензора от одного процесса (источника) всем процессам в группе. Эта операция является фундаментальной в распределённом обучении, так как позволяет синхронизировать начальные веса модели или отправлять обновления градиентов. Первым параметром функция принимает тензор для рассылки, вторым параметром указывается идентификатор процесса-источника.
Синтаксис
torch.distributed.broadcast(tensor, src, group=None, async_op=False)
Параметры
Функция принимает несколько параметров:
tensor - тензор, который будет отправлен из источника и получен всеми процессами. В процессе-источнике этот тензор должен быть заполнен данными, в остальных процессах он будет перезаписан полученными значениями.
src - целочисленный идентификатор процесса, который будет источником широковещательной рассылки. Только этот процесс отправляет свои данные, все остальные процессы получают их.
group - необязательный параметр, определяющий группу процессов, участвующих в операции. По умолчанию используется группа по умолчанию (все процессы).
async_op - необязательный булевый параметр. Если установлен в True, функция выполняется асинхронно и возвращает объект Work, который можно использовать для ожидания завершения операции.
Пример
Давайте выполним широковещательную рассылку от процесса с рангом 0 всем процессам:
import torch
import torch.distributed as dist
# Инициализация группы процессов
dist.init_process_group('gloo')
# Получаем ранг текущего процесса
rank = dist.get_rank()
# Создаем тензор с данными
if rank == 0:
t = torch.tensor([1, 2, 3, 4, 5])
else:
t = torch.tensor([0, 0, 0, 0, 0])
# Выполняем широковещательную рассылку от процесса с рангом 0
dist.broadcast(t, src=0)
print(f"Process {rank}: {t}")
# Завершаем работу группы
dist.destroy_process_group()
Результат выполнения кода для всех процессов:
Process 0: tensor([1, 2, 3, 4, 5])
Process 1: tensor([1, 2, 3, 4, 5])
Process 2: tensor([1, 2, 3, 4, 5])
Process 3: tensor([1, 2, 3, 4, 5])
Пример
Давайте используем асинхронную широковещательную рассылку с проверкой завершения:
import torch
import torch.distributed as dist
# Инициализация группы процессов
dist.init_process_group('gloo')
# Получаем ранг текущего процесса
rank = dist.get_rank()
# Создаем тензоры разных размеров
if rank == 0:
t = torch.tensor([10, 20, 30, 40, 50, 60])
else:
t = torch.tensor([0, 0, 0, 0, 0, 0])
# Асинхронная широковещательная рассылка
work = dist.broadcast(t, src=0, async_op=True)
# Выполняем другие вычисления во время ожидания
print(f"Process {rank}: doing other work...")
# Ожидаем завершения операции
work.wait()
print(f"Process {rank}: completed broadcast, result = {t}")
# Завершаем работу группы
dist.destroy_process_group()
Результат выполнения кода:
Process 0: doing other work...
Process 1: doing other work...
Process 0: completed broadcast, result = tensor([10, 20, 30, 40, 50, 60])
Process 1: completed broadcast, result = tensor([10, 20, 30, 40, 50, 60])
Пример
Давайте выполним широковещательную рассылку с использованием пользовательской группы процессов:
import torch
import torch.distributed as dist
# Инициализация группы процессов
dist.init_process_group('gloo')
# Получаем ранг и размер группы
rank = dist.get_rank()
world_size = dist.get_world_size()
# Создаем подгруппу из процессов с четными рангами
ranks = list(range(0, world_size, 2))
group = dist.new_group(ranks)
# Создаем тензор для рассылки
if rank in ranks:
if rank == 0:
t = torch.tensor([100, 200, 300])
else:
t = torch.tensor([0, 0, 0])
# Выполняем широковещательную рассылку в подгруппе
dist.broadcast(t, src=0, group=group)
print(f"Process {rank} (in subgroup): {t}")
else:
print(f"Process {rank} (not in subgroup): no broadcast")
# Завершаем работу группы
dist.destroy_process_group()
Результат выполнения кода:
Process 0 (in subgroup): tensor([100, 200, 300])
Process 1 (not in subgroup): no broadcast
Process 2 (in subgroup): tensor([100, 200, 300])
Process 3 (not in subgroup): no broadcast
Смотрите также
-
функцию
all_reduce,
которая выполняет операцию редукции и рассылки результата всем процессам -
функцию
all_gather,
которая собирает тензоры со всех процессов и рассылает результат каждому процессу -
функцию
init_process_group,
которая инициализирует группу процессов для распределённой работы -
функцию
get_rank,
которая возвращает ранг текущего процесса в группе