Spec-Zone.ru › PyTorch 2

torch.func.hessian

torch.func.hessian(func, argnums=0)

Вычисляет гессиан func по аргументу(ам) с индексом argnum с помощью стратегии «спереди-назад».

Стратегия «спереди-назад» (составляя jacfwd(jacrev(func))) — хороший дефолтный вариант для хорошей производительности. Можно вычислять гессианы с помощью других комбинаций jacfwd() и jacrev(), например jacfwd(jacfwd(func)) или jacrev(jacrev(func)).

Параметры
  • func (функция) – Python-функция, которая принимает один или несколько аргументов, один из которых должен быть тензором, и возвращает один или несколько тензоров
  • argnums (int или кортеж[int]) – необязательный целочисленный или кортеж целых чисел, указывающий, по отношению к каким аргументам вычислять гессиан. По умолчанию: 0.
Возвращает

Возвращает функцию, которая принимает такие же входные данные, как и func, и возвращает гессиан func по отношению к аргументу(ам) с индексом argnums.

Примечание

Вы можете столкнуться с ошибкой API «режим вычислений вперед не реализован для оператора X». В таком случае, пожалуйста, отправьте отчет об ошибке, и мы его рассмотрим. Альтернативой является использование jacrev(jacrev(func)), у которого лучше покрыты операторы.

Пример использования с функцией R^N -> R^1 даёт гессиан NxN:

>>> from torch.func import hessian
>>> def f(x):
>>>   return x.sin().sum()
>>>
>>> x = torch.randn(5)
>>> hess = hessian(f)(x)  # equivalent to jacfwd(jacrev(f))(x)
>>> assert torch.allclose(hess, torch.diag(-x.sin()))

© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/2.1/generated/torch.func.hessian.html

Spec-Zone.ru

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