Spec-Zone.ru › PyTorch 2

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 и функцию, которая вычисляет vjp func относительно всех 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

Spec-Zone.ru

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