Spec-Zone.ru › PyTorch 2

torch.func.linearize

torch.func.linearize(func, *primals)

Возвращает значение func в точке primals и линейное приближение в точке primals.

Параметры
  • func (Callable) – Функция Python, принимающая один или несколько аргументов.
  • primals (Tensors) – Позиционные аргументы для func, все из которых должны быть тензорами. Это значения, в которых функция линейно аппроксимируется.
Возвращает

Возвращает кортеж (output, jvp_fn), содержащий результат применения func к primals и функцию, вычисляющую jvp func в точке primals.

Тип возвращаемого значения

Tuple[Any, Callable]

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

Spec-Zone.ru

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