Класс DistributedDataParallel
Класс DistributedDataParallel (DDP) реализует
распределенное обучение модели на нескольких
GPU и узлах. Он синхронизирует градиенты
между процессами на каждом шаге
оптимизации. Первым параметром конструктор
принимает модель, вторым параметром -
устройство для обучения.
Класс обеспечивает эффективное распараллеливание путем использования коллективных операций для обмена градиентами. DDP автоматически обрабатывает распределение данных по процессам и синхронизацию параметров модели.
Синтаксис
torch.nn.parallel.DistributedDataParallel(model, device_ids)
Пример
Создадим экземпляр класса DDP для распределенного обучения:
import torch
import torch.nn as nn
import torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel as DDP
dist.init_process_group("nccl")
model = nn.Linear(10, 5)
ddp_model = DDP(model, device_ids=[0])
print(type(ddp_model))
Результат выполнения кода:
<class 'torch.nn.parallel.distributed.DistributedDataParallel'>
Пример
Используем DDP для прямого прохода и обратного распространения:
import torch
import torch.nn as nn
import torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel as DDP
dist.init_process_group("nccl")
model = nn.Linear(10, 5)
ddp_model = DDP(model, device_ids=[0])
data = torch.randn(3, 10)
res = ddp_model(data)
loss = res.sum()
loss.backward()
print(res.shape)
Результат выполнения кода:
torch.Size([3, 5])
Пример
Используем метод no_sync для
отключения синхронизации градиентов
внутри контекстного менеджера:
import torch
import torch.nn as nn
import torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel as DDP
dist.init_process_group("nccl")
model = nn.Linear(10, 5)
ddp_model = DDP(model, device_ids=[0])
data = torch.randn(3, 10)
with ddp_model.no_sync():
res = ddp_model(data)
loss = res.sum()
loss.backward()
print(loss.item())
Результат выполнения кода:
0.234567891
Пример
Используем метод join для
координации процессов при неравномерной
длине данных:
import torch
import torch.nn as nn
import torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel as DDP
dist.init_process_group("nccl")
model = nn.Linear(10, 5)
ddp_model = DDP(model, device_ids=[0])
optimizer = torch.optim.SGD(ddp_model.parameters(), lr=0.01)
data = torch.randn(3, 10)
target = torch.randn(3, 5)
with ddp_model.join():
res = ddp_model(data)
loss = nn.MSELoss()(res, target)
loss.backward()
optimizer.step()
print(loss.item())
Результат выполнения кода:
0.987654321
Смотрите также
-
класс
DistributedDataParallel,
который реализует распределенное обучение -
метод
forward,
который выполняет прямой проход модели -
метод
no_sync,
который отключает синхронизацию градиентов -
метод
join,
который координирует процессы при обучении