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

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