torch.func.linearize
-
torch.func.linearize(func, *primals)[исходный код] -
Возвращает значение
funcприprimalsи линейное приближение приprimals.- Параметры:
-
- func (Callable[..., Any]) – Функция Python, принимающая один или несколько аргументов.
-
primals (Tensors) – Позиционные аргументы функции
func, каждый из которых должен быть тензором. Это значения, в которых вычисляется линейное приближение функции.
- Возвращает:
-
Возвращает кортеж
(output, jvp_fn), содержащий результат примененияfuncкprimalsи функцию, вычисляющую jvp дляfunc, вычисленной вprimals. - Тип возвращаемого значения:
-
tuple[Any, Callable[…, Any]]
linearize полезна, если jvp нужно вычислить несколько раз в
primals. Однако для этого linearize сохраняет промежуточные результаты вычислений и требует больше памяти, чем непосредственное применение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.]]) >>>
© 2026, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://docs.pytorch.org/docs/2.14/generated/torch.func.linearize.html