Функция all_reduce
Функция distributed.all_reduce выполняет коллективную операцию редукции данных
между всеми процессами в группе. Эта операция суммирует, перемножает или выполняет
другую агрегацию значений тензоров от всех процессов и записывает результат обратно
в тот же тензор каждого процесса.
Первый параметр функции - это тензор tensor, который будет изменён.
Второй параметр - операция редукции op (по умолчанию ReduceOp.SUM).
Третий параметр - группа процессов group, в которой выполняется операция.
Функция работает синхронно и блокирует выполнение до завершения операции.
Синтаксис
import torch.distributed as dist
dist.all_reduce(tensor, op=ReduceOp.SUM, group=None)
Пример
Базовый пример использования all_reduce для суммирования тензоров
между двумя процессами:
Результат выполнения кода для двух процессов:
Process 0 before all_reduce: tensor([1., 2.])
Process 1 before all_reduce: tensor([2., 3.])
Process 0 after all_reduce: tensor([3., 5.])
Process 1 after all_reduce: tensor([3., 5.])
Пример
Операция all_reduce с умножением:
Результат выполнения кода для двух процессов:
Process 0 before all_reduce: tensor([1., 2.])
Process 1 before all_reduce: tensor([2., 3.])
Process 0 after all_reduce: tensor([2., 6.])
Process 1 after all_reduce: tensor([2., 6.])
Пример
Использование all_reduce с пользовательской группой процессов
и операцией нахождения максимума:
Результат выполнения кода для трёх процессов:
Process 0 before all_reduce: tensor([1., 2.])
Process 2 before all_reduce: tensor([3., 4.])
Process 0 after all_reduce: tensor([3., 4.])
Process 2 after all_reduce: tensor([3., 4.])
Пример
Использование all_reduce для вычисления средней точности модели
на всех устройствах при распределённом обучении:
Результат выполнения кода для четырёх процессов:
Process 0 local accuracy: 0.83
Process 1 local accuracy: 0.95
Process 2 local accuracy: 0.71
Process 3 local accuracy: 0.88
Process 0 global average accuracy: 0.84
Process 1 global average accuracy: 0.84
Process 2 global average accuracy: 0.84
Process 3 global average accuracy: 0.84
Смотрите также
-
функцию
init_process_group,
которая инициализирует распределённую среду -
функцию
all_gather,
которая собирает тензоры со всех процессов -
функцию
broadcast,
которая отправляет тензор от одного процесса всем остальным -
функцию
reduce,
которая выполняет редукцию данных на одном процессе