Spec-Zone.ru › PyTorch 2

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.

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

Any

© 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

Spec-Zone.ru

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