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

Метод __call__

Метод __call__ класса Module делает экземпляр модели вызываемым, подобно функции. При вызове модели, переданные входные данные направляются в метод forward, но при этом автоматически обрабатываются хуки, режимы обучения и градиенты. Этот метод не рекомендуется переопределять - вместо этого следует реализовывать логику в forward.

Синтаксис

# Вызов метода __call__ через экземпляр output = model(input_tensor)

Параметры:

  • *args, **kwargs - произвольные аргументы, передаваемые в метод forward;
  • возвращаемое значение - результат вызова forward.

Пример простейшего вызова

Создадим линейный слой и вызовем его как функцию:

import torch layer = torch.nn.Linear(3, 2) t = torch.randn(1, 3) res = layer(t) print(res)

Результат выполнения кода:

tensor([[ 0.3193, -0.0556]], grad_fn=<AddmmBackward0>)

Обратите внимание, что у тензора появился атрибут grad_fn, что говорит о включенном механизме автоградиента.

Пример с пользовательским модулем

Создадим свою модель и вызовем её, передав несколько аргументов:

import torch class MyModel(torch.nn.Module): def __init__(self): super().__init__() self.linear = torch.nn.Linear(2, 1) def forward(self, x, add_const=0.0): res = self.linear(x) return res + add_const model = MyModel() t = torch.tensor([[1.0, 2.0]]) res = model(t, add_const=5.0) print(res)

Результат выполнения кода:

tensor([[5.6029]], grad_fn=<AddBackward0>)

Метод __call__ передал дополнительный аргумент add_const в forward.

Особенности работы с градиентами

При вызове модели через __call__ градиенты вычисляются автоматически, если тензоры требуют их. В отличие от прямого вызова forward, такой подход гарантирует правильную работу хуков и режимов обучения:

import torch torch.manual_seed(0) model = torch.nn.Linear(2, 1) x = torch.randn(1, 2, requires_grad=True) # Вызов через __call__ res = model(x) res.backward() print(x.grad)

Результат выполнения кода:

tensor([[-0.0930, 0.6668]])

Если бы мы вызвали forward напрямую, градиенты не были бы автоматически связаны с параметрами модели.

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

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