Метод no_sync
Метод no_sync класса DistributedDataParallel
возвращает контекстный менеджер, который отключает
синхронизацию градиентов между процессами во время
выполнения обратного распространения. Это полезно
для накопления градиентов на нескольких мини-батчах
перед выполнением одного шага оптимизации. Метод
не принимает никаких параметров.
Синтаксис
with model.no_sync():
output = model(input)
loss = criterion(output, target)
loss.backward()
Пример
Давайте создадим модель и применим метод 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(backend='nccl')
model = nn.Linear(10, 5).cuda()
ddp_model = DDP(model)
optimizer = torch.optim.SGD(ddp_model.parameters(), lr=0.01)
for i in range(2):
with ddp_model.no_sync():
input = torch.randn(4, 10).cuda()
target = torch.randn(4, 5).cuda()
output = ddp_model(input)
loss = nn.MSELoss()(output, target)
loss.backward()
optimizer.step()
optimizer.zero_grad()
Пример
Давайте сравним поведение с синхронизацией и без неё при накоплении градиентов:
import torch
import torch.nn as nn
import torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel as DDP
dist.init_process_group(backend='nccl')
model = nn.Linear(10, 5).cuda()
ddp_model = DDP(model)
# Накопление градиентов без синхронизации (эффективно)
for i in range(4):
with ddp_model.no_sync():
input = torch.randn(2, 10).cuda()
target = torch.randn(2, 5).cuda()
output = ddp_model(input)
loss = nn.MSELoss()(output, target)
loss.backward()
# Последний батч с синхронизацией
input = torch.randn(2, 10).cuda()
target = torch.randn(2, 5).cuda()
output = ddp_model(input)
loss = nn.MSELoss()(output, target)
loss.backward()
optimizer = torch.optim.SGD(ddp_model.parameters(), lr=0.01)
optimizer.step()
В этом примере первые четыре батча выполняются без синхронизации градиентов, а последний запускает синхронизацию и обновление весов.
Пример
Давайте используем метод 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(backend='nccl')
model = nn.Sequential(
nn.Linear(20, 64),
nn.ReLU(),
nn.Linear(64, 10)
).cuda()
ddp_model = DDP(model)
optimizer = torch.optim.Adam(ddp_model.parameters(), lr=0.001)
accumulation_steps = 4
for epoch in range(2):
for batch_idx in range(8):
input = torch.randn(16, 20).cuda()
target = torch.randint(0, 10, (16,)).cuda()
is_accumulating = (batch_idx + 1) % accumulation_steps != 0
if is_accumulating:
with ddp_model.no_sync():
output = ddp_model(input)
loss = nn.CrossEntropyLoss()(output, target)
loss.backward()
else:
output = ddp_model(input)
loss = nn.CrossEntropyLoss()(output, target)
loss.backward()
optimizer.step()
optimizer.zero_grad()
Результат выполнения кода:
Градиенты накапливаются в течение 4 батчей,
затем выполняется синхронизация и шаг оптимизации
Смотрите также
-
класс
DistributedDataParallel,
который обеспечивает распределённое обучение с синхронизацией градиентов -
метод
forward,
который выполняет прямой проход в распределённой модели -
метод
join,
который синхронизирует процессы при неравномерных батчах -
метод
no_sync,
который отключает синхронизацию градиентов для накопления