Spec-Zone.ru › PyTorch 2.14

torch.autograd.Function.forward

static Function.forward(*args, **kwargs) [исходный код]

Определите прямой проход пользовательской функции autograd.

Эту функцию необходимо переопределить во всех подклассах. Есть два способа определить forward:

Использование 1 (объединённые forward и ctx):

@staticmethod
def forward(ctx: Any, *args: Any, **kwargs: Any) -> Any:
    pass
  • Первым аргументом функция должна принимать контекст ctx, за которым может следовать любое количество аргументов (тензоров или объектов других типов).
  • Подробнее см. в разделе Объединённые или раздельные forward() и setup_context().

Использование 2 (раздельные forward и ctx):

@staticmethod
def forward(*args: Any, **kwargs: Any) -> Any:
    pass


@staticmethod
def setup_context(ctx: Any, inputs: Tuple[Any, ...], output: Any) -> None:
    pass
  • Функция forward больше не принимает аргумент ctx.
  • Вместо этого необходимо также переопределить статический метод torch.autograd.Function.setup_context() для настройки объекта ctx. output — это результат forward, а inputs — кортеж входных данных forward.
  • Подробнее см. в разделе Расширение torch.autograd.

Контекст можно использовать для хранения произвольных данных, которые затем можно получить во время обратного прохода. Тензоры не следует непосредственно сохранять в ctx (хотя это пока не запрещено ради обратной совместимости). Вместо этого тензоры следует сохранять с помощью ctx.save_for_backward(), если они предназначены для использования в backward (то есть в vjp), или с помощью ctx.save_for_forward(), если они предназначены для использования в jvp.

Тип возвращаемого значения:

Any

© 2026, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://docs.pytorch.org/docs/2.14/generated/torch.autograd.Function.forward.html

Spec-Zone.ru

Настройки Оффлайн Что нового Помощь О нас
Spec-Zone .ru
спецификации, руководства, описания, API