torch.func.vjp
-
torch.func.vjp(func, *primals, has_aux=False) -
Возвращает кортеж, содержащий результаты
funcприменительно кprimals, и функцию, которая, при полученииcotangents, вычисляет обратный якобианfuncотносительноprimalsумноженный наcotangents.- Параметры
-
- func (Callable) – Функция Python, принимающая один или несколько аргументов. Должна возвращать один или несколько тензоров.
-
primals (Tensors) – Позиционные аргументы для
func, которые должны быть тензорами. Возвращаемая функция также будет вычислять производную относительно этих аргументов. -
has_aux (bool) – Флаг, указывающий, что
funcвозвращает кортеж(output, aux), где первый элемент — результат функции, подлежащей дифференцированию, а второй — другие вспомогательные объекты, которые не будут дифференцироваться. По умолчанию: False.
- Возвращаемое значение
-
Возвращает кортеж
(output, vjp_fn), содержащий результатfuncпо отношению кprimalsи функцию, которая вычисляет vjpfuncотносительно всехprimalsс использованием котангенсов, переданных в возвращаемую функцию. Еслиhas_aux is True, то вместо этого возвращается кортеж(output, vjp_fn, aux). Возвращаемая функцияvjp_fnбудет возвращать кортеж из каждого VJP.
В простых случаях
vjp()ведет себя так же, какgrad()>>> x = torch.randn([5]) >>> f = lambda x: x.sin().sum() >>> (_, vjpfunc) = torch.func.vjp(f, x) >>> grad = vjpfunc(torch.tensor(1.))[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. Все ключевые слова используют своё значение по умолчанию.>>> 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.))
Примечание
Использование 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.
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/2.1/generated/torch.func.vjp.html