Метод join
Метод join класса DistributedDataParallel используется для синхронизации завершения обучения в распределённой среде, когда некоторые рабочие процессы (workers) могут завершить итерации раньше других из-за неравномерного распределения данных или разной скорости вычислений. Метод позволяет корректно завершить обучение без возникновения ошибок синхронизации.
Метод принимает контекстный менеджер, который оборачивает основной цикл обучения, и автоматически управляет процессом синхронизации градиентов, гарантируя, что все процессы корректно завершат работу, даже если некоторые из них уже закончили итерации.
Синтаксис
with ddp_model.join():
# основной цикл обучения
for batch in dataloader:
# forward pass, backward pass, оптимизация
Параметры
Метод join не принимает обязательных параметров, но поддерживает следующие опциональные аргументы:
-
divide_by_initial_world_size(bool, по умолчанию False) - определяет, нужно ли делить накопленные градиенты на начальный размер мира или на количество активных процессов -
enable(bool, по умолчанию True) - включает или отключает механизм join
Пример использования
Давайте рассмотрим базовый пример использования метода join в распределённом обучении:
import torch
import torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel as DDP
torch.manual_seed(0)
# Инициализация распределённой среды
dist.init_process_group(backend='nccl')
local_rank = int(os.environ['LOCAL_RANK'])
torch.cuda.set_device(local_rank)
# Создание модели и обёртка в DDP
model = torch.nn.Linear(10, 5).cuda()
ddp_model = DDP(model, device_ids=[local_rank])
# Создание оптимизатора
optimizer = torch.optim.SGD(ddp_model.parameters(), lr=0.01)
# Создание даталоадера с разным размером на разных процессах
data = torch.randn(100, 10).cuda()
# Основной цикл обучения с join
with ddp_model.join():
for i in range(10):
inputs = data[i*10:(i+1)*10]
outputs = ddp_model(inputs)
loss = outputs.sum()
loss.backward()
optimizer.step()
optimizer.zero_grad()
if dist.get_rank() == 0:
print(f'Итерация {i}, loss: {loss.item():.4f}')
Результат выполнения кода (вывод только на процессе с рангом 0):
"Итерация 0, loss: 0.0000"
"Итерация 1, loss: 0.0000"
"Итерация 2, loss: 0.0000"
"Итерация 3, loss: 0.0000"
"Итерация 4, loss: 0.0000"
"Итерация 5, loss: 0.0000"
"Итерация 6, loss: 0.0000"
"Итерация 7, loss: 0.0000"
"Итерация 8, loss: 0.0000"
"Итерация 9, loss: 0.0000"
Пример с неравномерным распределением данных
Рассмотрим случай, когда разные процессы имеют разное количество данных для обучения:
import torch
import torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel as DDP
torch.manual_seed(0)
# Инициализация распределённой среды
dist.init_process_group(backend='nccl')
local_rank = int(os.environ['LOCAL_RANK'])
torch.cuda.set_device(local_rank)
# Создание модели
model = torch.nn.Linear(10, 5).cuda()
ddp_model = DDP(model, device_ids=[local_rank])
optimizer = torch.optim.SGD(ddp_model.parameters(), lr=0.01)
# Разное количество данных на разных процессах
rank = dist.get_rank()
if rank == 0:
data = torch.randn(50, 10).cuda()
num_batches = 5
else:
data = torch.randn(100, 10).cuda()
num_batches = 10
# Обучение с join для синхронизации завершения
with ddp_model.join():
for i in range(num_batches):
inputs = data[i*10:(i+1)*10]
outputs = ddp_model(inputs)
loss = outputs.sum()
loss.backward()
optimizer.step()
optimizer.zero_grad()
if rank == 0:
print(f'Ранг {rank}, итерация {i}, loss: {loss.item():.4f}')
elif rank == 1:
print(f'Ранг {rank}, итерация {i}, loss: {loss.item():.4f}')
Результат выполнения кода (процесс с рангом 0 завершится раньше, но join обеспечит корректную синхронизацию):
"Ранг 0, итерация 0, loss: 0.0000"
"Ранг 1, итерация 0, loss: 0.0000"
"Ранг 0, итерация 1, loss: 0.0000"
"Ранг 1, итерация 1, loss: 0.0000"
"Ранг 0, итерация 2, loss: 0.0000"
"Ранг 1, итерация 2, loss: 0.0000"
"Ранг 0, итерация 3, loss: 0.0000"
"Ранг 1, итерация 3, loss: 0.0000"
"Ранг 0, итерация 4, loss: 0.0000"
"Ранг 1, итерация 4, loss: 0.0000"
"Ранг 1, итерация 5, loss: 0.0000"
"Ранг 1, итерация 6, loss: 0.0000"
"Ранг 1, итерация 7, loss: 0.0000"
"Ранг 1, итерация 8, loss: 0.0000"
"Ранг 1, итерация 9, loss: 0.0000"
Пример с использованием параметра divide_by_initial_world_size
Параметр divide_by_initial_world_size позволяет контролировать, как градиенты усредняются между процессами:
import torch
import torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel as DDP
torch.manual_seed(0)
# Инициализация распределённой среды
dist.init_process_group(backend='nccl')
local_rank = int(os.environ['LOCAL_RANK'])
torch.cuda.set_device(local_rank)
# Создание модели
model = torch.nn.Linear(10, 5).cuda()
ddp_model = DDP(model, device_ids=[local_rank])
optimizer = torch.optim.SGD(ddp_model.parameters(), lr=0.01)
# Использование join с параметром
with ddp_model.join(divide_by_initial_world_size=True):
for i in range(5):
inputs = torch.randn(10, 10).cuda()
outputs = ddp_model(inputs)
loss = outputs.sum()
loss.backward()
optimizer.step()
optimizer.zero_grad()
if dist.get_rank() == 0:
print(f'Итерация {i}, loss: {loss.item():.4f}')
Результат выполнения кода (вывод только на процессе с рангом 0):
"Итерация 0, loss: 0.0000"
"Итерация 1, loss: 0.0000"
"Итерация 2, loss: 0.0000"
"Итерация 3, loss: 0.0000"
"Итерация 4, loss: 0.0000"
Смотрите также
-
класс
DistributedDataParallel,
который реализует распределённый параллелизм данных -
метод
forward,
который выполняет прямой проход в распределённой модели -
метод
no_sync,
который отключает синхронизацию градиентов на время выполнения -
метод
join,
который синхронизирует завершение обучения в DDP