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.- Тип возвращаемого значения:
© 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