РЕПЕТИТОР математика физика информатика
Для школьников и студентов. Подтягивание пробелов. ЦЭ, ЦТ, ОГЭ, ЕГЭ.
Идет набор на ЛЕТО. Жмите для подробностей:)
738 of 769 menu

Метод 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,
    который отключает синхронизацию градиентов для накопления
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить