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