torch.func.vjp
-
torch.func.vjp(func, *primals, has_aux=False)[source] -
Сокращение от vector-Jacobian product (произведение вектора на матрицу Якоби): возвращает кортеж, содержащий результаты применения
funcкprimals, и функцию, которая при полученииcotangentsвычисляет якобианfuncв режиме обратного распространения по отношению кprimals, умноженный наcotangents.- Параметры:
-
- func (Callable[..., Any]) – Функция Python, принимающая один или несколько аргументов. Должна возвращать один или несколько тензоров.
-
primals (Tensors) – Позиционные аргументы для
func, каждый из которых должен быть тензором. Возвращаемая функция также вычисляет производную по этим аргументам -
has_aux (bool) – Флаг, указывающий, что
funcвозвращает кортеж(output, aux), в котором первый элемент — результат функции, для которой вычисляется производная, а второй — другие вспомогательные объекты, для которых производная вычисляться не будет. Значение по умолчанию: False.
- Возвращает:
-
Возвращает кортеж
(output, vjp_fn), содержащий результат примененияfuncкprimals, и функцию, которая вычисляет vjp дляfuncпо всемprimals, используя котангенты, переданные возвращаемой функции. Еслиhas_aux is True, возвращается кортеж(output, vjp_fn, aux). Возвращаемая функцияvjp_fnвозвращает кортеж, содержащий каждый VJP. - Тип возвращаемого значения:
-
tuple[Any, Callable[…, Any]] | tuple[Any, Callable[…, Any], Any]
В простых случаях
vjp()работает так же, какgrad()>>> x = torch.randn([5]) >>> f = lambda x: x.sin().sum() >>> (_, vjpfunc) = torch.func.vjp(f, x) >>> grad = vjpfunc(torch.tensor(1.0))[0] >>> assert torch.allclose(grad, torch.func.grad(f)(x))
Однако
vjp()поддерживает функции с несколькими выходами, если передать котангенты для каждого из выходов>>> x = torch.randn([5]) >>> f = lambda x: (x.sin(), x.cos()) >>> (_, vjpfunc) = torch.func.vjp(f, x) >>> vjps = vjpfunc((torch.ones([5]), torch.ones([5]))) >>> assert torch.allclose(vjps[0], x.cos() + -x.sin())
vjp()также поддерживает выходные данные в виде структур Python>>> x = torch.randn([5]) >>> f = lambda x: {"first": x.sin(), "second": x.cos()} >>> (_, vjpfunc) = torch.func.vjp(f, x) >>> cotangents = {"first": torch.ones([5]), "second": torch.ones([5])} >>> vjps = vjpfunc(cotangents) >>> assert torch.allclose(vjps[0], x.cos() + -x.sin())Функция, возвращаемая
vjp(), вычисляет частные производные по каждому изprimals>>> x, y = torch.randn([5, 4]), torch.randn([4, 5]) >>> (_, vjpfunc) = torch.func.vjp(torch.matmul, x, y) >>> cotangents = torch.randn([5, 5]) >>> vjps = vjpfunc(cotangents) >>> assert len(vjps) == 2 >>> assert torch.allclose(vjps[0], torch.matmul(cotangents, y.transpose(0, 1))) >>> assert torch.allclose(vjps[1], torch.matmul(x.transpose(0, 1), cotangents))
primals— это позиционные аргументы дляf. Для всех kwargs используются значения по умолчанию>>> x = torch.randn([5]) >>> def f(x, scale=4.): >>> return x * scale >>> >>> (_, vjpfunc) = torch.func.vjp(f, x) >>> vjps = vjpfunc(torch.ones_like(x)) >>> assert torch.allclose(vjps[0], torch.full(x.shape, 4.0))
Примечание
Использование PyTorch
torch.no_gradвместе сvjp. Случай 1: использованиеtorch.no_gradвнутри функции:>>> def f(x): >>> with torch.no_grad(): >>> c = x ** 2 >>> return x - c
В этом случае
vjp(f)(x)будет учитывать внутреннийtorch.no_grad.Случай 2: использование
vjpвнутри менеджера контекстаtorch.no_grad:>>> with torch.no_grad(): >>> vjp(f)(x)
В этом случае
vjpбудет учитывать внутреннийtorch.no_grad, но не внешний. Это связано с тем, чтоvjp— это «преобразование функции»: результат не должен зависеть от результата менеджера контекста, находящегося за пределамиf.
© 2026, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://docs.pytorch.org/docs/2.14/generated/torch.func.vjp.html