Метод forward
Метод forward класса DataParallel выполняет прямое распространение
входных данных через копии модели, размещённые на нескольких устройствах
(GPU). В качестве первого параметра метод принимает входной тензор
или кортеж тензоров. Метод автоматически распределяет входные данные
по устройствам, выполняет прямой проход на каждом устройстве, собирает
результаты и возвращает объединённый выходной тензор.
Синтаксис
output = data_parallel_model.forward(*inputs, **kwargs)
Метод вызывается автоматически при передаче данных в объект
DataParallel. Входные данные распределяются по пакетному
измерению (dimension 0) между доступными устройствами.
Пример
Создадим простую модель и обернём её в DataParallel для
использования двух GPU:
import torch
import torch.nn as nn
# Простая модель
class SimpleModel(nn.Module):
def __init__(self):
super(SimpleModel, self).__init__()
self.fc = nn.Linear(10, 5)
def forward(self, x):
return self.fc(x)
# Создаём модель и оборачиваем в DataParallel
model = SimpleModel()
if torch.cuda.device_count() > 1:
model = nn.DataParallel(model)
# Создаём входные данные
x = torch.randn(4, 10)
# Вызов forward происходит автоматически
output = model(x)
print(output.shape)
Результат выполнения кода:
torch.Size([4, 5])
Пример
Метод forward поддерживает передачу нескольких аргументов.
Рассмотрим пример с моделью, принимающей два тензора:
import torch
import torch.nn as nn
class MultiInputModel(nn.Module):
def __init__(self):
super(MultiInputModel, self).__init__()
self.fc1 = nn.Linear(5, 3)
self.fc2 = nn.Linear(3, 2)
def forward(self, x1, x2):
h = self.fc1(x1 + x2)
return self.fc2(h)
model = MultiInputModel()
if torch.cuda.device_count() > 1:
model = nn.DataParallel(model)
t1 = torch.randn(4, 5)
t2 = torch.randn(4, 5)
output = model(t1, t2)
print(output.shape)
Результат выполнения кода:
torch.Size([4, 2])
Пример
Метод forward также работает с именованными аргументами.
В этом случае они передаются во все копии модели:
import torch
import torch.nn as nn
class KwargsModel(nn.Module):
def __init__(self):
super(KwargsModel, self).__init__()
self.fc = nn.Linear(10, 5)
def forward(self, x, scale=1.0):
return self.fc(x) * scale
model = KwargsModel()
if torch.cuda.device_count() > 1:
model = nn.DataParallel(model)
x = torch.randn(4, 10)
output = model(x, scale=2.0)
print(output.shape)
Результат выполнения кода:
torch.Size([4, 5])
Смотрите также
-
класс
DataParallel,
который реализует параллельное выполнение модели на нескольких GPU -
метод
forward,
который выполняет прямой проход в модулях PyTorch -
класс
DistributedDataParallel,
который обеспечивает распределённое обучение по нескольким узлам -
метод
to,
который перемещает модель или тензор на указанное устройство