torch.autograd.Function.jvp
-
static Function.jvp(ctx, *grad_inputs) -
Определяет формулу для дифференцирования операции с автоматическим дифференцированием в прямом режиме. Эта функция должна быть переопределена всеми подклассами. Она должна принимать контекст
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/2.1/generated/torch.autograd.Function.jvp.html