torch.autograd.Function.jvp
-
static Function.jvp(ctx, *grad_inputs)[source] -
Определяет формулу для дифференцирования операции с помощью автоматического дифференцирования в прямом режиме. Эта функция должна быть переопределена всеми подклассами. Она должна принимать контекст
ctxв качестве первого аргумента, а также столько же входных данных, сколько и вforward()(для нетензорных входных данных функции forward будет передано значение None), и она должна возвращать столько же тензоров, сколько выходов уforward(). Каждый аргумент — градиент по отношению к соответствующему входу, а каждое возвращаемое значение — градиент по отношению к соответствующему выходу. Если вывод не является тензором или функция не дифференцируема по этому выходу, вы можете просто передать None в качестве градиента для этого входа.Вы можете использовать объект
ctxдля передачи любого значения из forward в эту функцию.- Тип возвращаемого значения:
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/1.13/generated/torch.autograd.Function.jvp.html