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

Метод 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])"

Смотрите также

  • класс Module,
    базовый класс для всех нейросетевых модулей
  • метод __call__,
    который вызывает forward при вызове объекта
  • метод train,
    который переключает модуль в режим обучения
  • метод eval,
    который переключает модуль в режим оценки
Мы используем cookie для работы сайта, аналитики и персонализации. Обработка данных происходит согласно Политике конфиденциальности.
принять все настроить отклонить