Spec-Zone.ru › PyTorch 2.14

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

Spec-Zone.ru

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