torch.autograd.Function.forward
-
static Function.forward(ctx, *args, **kwargs) -
Этот метод должен быть переопределён всеми подклассами. Существует два способа определения 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.- Тип возвращаемого значения
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/2.1/generated/torch.autograd.Function.forward.html