torch.func.linearize
-
torch.func.linearize(func, *primals) -
Возвращает значение
funcв точкеprimalsи линейное приближение в точкеprimals.- Параметры
-
- func (Callable) – Функция Python, принимающая один или несколько аргументов.
-
primals (Tensors) – Позиционные аргументы для
func, все из которых должны быть тензорами. Это значения, в которых функция линейно аппроксимируется.
- Возвращает
-
Возвращает кортеж
(output, jvp_fn), содержащий результат примененияfuncкprimalsи функцию, вычисляющую jvpfuncв точкеprimals. - Тип возвращаемого значения
linearize полезна, если jvp необходимо вычислить многократно в
primals. Однако для этого линейная аппроксимация сохраняет промежуточные вычисления и имеет более высокие требования к памяти, чем прямое применениеjvp. Поэтому, если всеtangentsизвестны, может быть более эффективно вычислить vmap(jvp) вместо использования linearize.Примечание
linearize вычисляет
funcдважды. Пожалуйста, отправьте запрос на реализацию с одним вычислением.- Пример::
-
>>> import torch >>> from torch.func import linearize >>> def fn(x): ... return x.sin() ... >>> output, jvp_fn = linearize(fn, torch.zeros(3, 3)) >>> jvp_fn(torch.ones(3, 3)) tensor([[1., 1., 1.], [1., 1., 1.], [1., 1., 1.]]) >>>
© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/2.1/generated/torch.func.linearize.html