Метод forward
Метод forward класса Module определяет,
как входные данные преобразуются в выходные
при прямом проходе через слой или модель.
Этот метод должен быть переопределён в
пользовательских классах, наследующих от
Module. Первым параметром метод принимает
входной тензор x, а возвращает тензор
результата.
Синтаксис
class MyModule(nn.Module):
def forward(self, x):
# Логика преобразования данных
return x
Пример
Давайте создадим простой линейный слой
с переопределённым методом forward:
import torch
import torch.nn as nn
class LinearLayer(nn.Module):
def __init__(self, in_features, out_features):
super().__init__()
self.weight = nn.Parameter(torch.randn(in_features, out_features))
self.bias = nn.Parameter(torch.randn(out_features))
def forward(self, x):
return x @ self.weight + self.bias
layer = LinearLayer(3, 2)
x = torch.randn(1, 3)
res = layer.forward(x)
print(res)
Результат выполнения кода:
tensor([[-2.1844, 0.3042]])
Пример
Переопределим метод forward для создания
блока с активацией ReLU:
import torch
import torch.nn as nn
class Block(nn.Module):
def __init__(self, in_features, out_features):
super().__init__()
self.linear = nn.Linear(in_features, out_features)
self.relu = nn.ReLU()
def forward(self, x):
x = self.linear(x)
x = self.relu(x)
return x
block = Block(4, 2)
x = torch.tensor([[1.0, -2.0, 3.0, -4.0]])
res = block.forward(x)
print(res)
Результат выполнения кода:
tensor([[0.0000, 0.0000]])
Пример
Метод forward может принимать несколько
входных аргументов для сложных моделей:
import torch
import torch.nn as nn
class SumModule(nn.Module):
def forward(self, x1, x2):
return x1 + x2
module = SumModule()
a = torch.tensor([1.0, 2.0, 3.0])
b = torch.tensor([4.0, 5.0, 6.0])
res = module.forward(a, b)
print(res)
Результат выполнения кода:
tensor([5., 7., 9.])
Пример
Метод forward автоматически вызывается
при вызове объекта как функции благодаря
методу __call__:
import torch
import torch.nn as nn
class SimpleModule(nn.Module):
def forward(self, x):
return x * 2
module = SimpleModule()
x = torch.tensor([1, 2, 3, 4, 5])
# Вызов через forward
res1 = module.forward(x)
# Вызов через объект как функцию
res2 = module(x)
print("Через forward:", res1)
print("Через объект:", res2)
Результат выполнения кода:
"Через forward: tensor([2, 4, 6, 8, 10])"
"Через объект: tensor([2, 4, 6, 8, 10])"