Метод forward
Метод forward класса DistributedDataParallel (DDP) выполняет прямое распространение
входных данных через распределенную модель. В отличие от стандартного вызова модели,
этот метод автоматически синхронизирует градиенты между процессами при обратном распространении.
Первый параметр метода - входной тензор или кортеж тензоров, которые передаются в основную модель.
Метод возвращает результат прямого распространения, идентичный результату обычной модели.
Синтаксис
output = ddp_model.forward(*inputs, **kwargs)
Параметры
Метод принимает следующие аргументы:
-
*inputs- позиционные аргументы, передаваемые в основную модель (тензоры или их кортеж); -
**kwargs- именованные аргументы, передаваемые в основную модель.
Возвращаемое значение
Метод возвращает результат прямого распространения основной модели - тензор или кортеж тензоров.
Пример
Создадим распределенную модель и выполним прямое распространение на одном процессе:
import torch
import torch.nn as nn
from torch.nn.parallel import DistributedDataParallel as DDP
# Инициализация распределенной среды (для примера)
# torch.distributed.init_process_group("gloo", rank=0, world_size=1)
model = nn.Linear(10, 5)
ddp_model = DDP(model)
x = torch.randn(3, 10)
output = ddp_model.forward(x)
print(output.shape)
Результат выполнения кода:
torch.Size([3, 5])
Пример
Передача нескольких аргументов в метод forward:
import torch
import torch.nn as nn
from torch.nn.parallel import DistributedDataParallel as DDP
class MyModel(nn.Module):
def forward(self, x, y):
return x + y
model = MyModel()
ddp_model = DDP(model)
a = torch.tensor([1.0, 2.0])
b = torch.tensor([3.0, 4.0])
res = ddp_model.forward(a, b)
print(res)
Результат выполнения кода:
tensor([4., 6.])
Пример
Использование именованных аргументов при вызове метода:
import torch
import torch.nn as nn
from torch.nn.parallel import DistributedDataParallel as DDP
class MyModel(nn.Module):
def forward(self, x, scale=1.0):
return x * scale
model = MyModel()
ddp_model = DDP(model)
x = torch.tensor([1.0, 2.0, 3.0])
res = ddp_model.forward(x, scale=2.0)
print(res)
Результат выполнения кода:
tensor([2., 4., 6.])
Пример
Прямое распространение с автоматической синхронизацией градиентов в распределенной среде:
import torch
import torch.nn as nn
from torch.nn.parallel import DistributedDataParallel as DDP
# torch.distributed.init_process_group("gloo", rank=0, world_size=1)
model = nn.Linear(5, 2)
ddp_model = DDP(model)
x = torch.randn(2, 5, requires_grad=True)
output = ddp_model.forward(x)
loss = output.sum()
loss.backward()
print(output)
print(model.weight.grad is not None)
Результат выполнения кода:
tensor([[...]], grad_fn=<AddmmBackward0>)
True
Смотрите также
-
класс
DistributedDataParallel,
который реализует распределенное обучение с синхронизацией градиентов -
метод
no_sync,
который отключает синхронизацию градиентов в контекстном менеджере -
метод
join,
который обеспечивает корректную работу с неоднородными входными данными -
класс
DistributedDataParallel,
который используется для распределенного обучения моделей