Spec-Zone.ru › PyTorch 2

torch.linalg.tensorinv

torch.linalg.tensorinv(A, ind=2, *, out=None) → Tensor

Вычисляет мультипликативную обратную величину torch.tensordot().

Если m является произведением первых ind измерений A и n является произведением остальных измерений, эта функция ожидает, что m и n будут равны. Если это так, она вычисляет тензор X такой, что tensordot(A, X, ind) является единичной матрицей в измерении m. X будет иметь форму A, но с первыми ind измерениями, переместинными в конец.

X.shape == A.shape[ind:] + A.shape[:ind]

Поддерживает входные типы float, double, cfloat и cdouble.

Примечание

Когда A является тензором размерности 2 и ind= 1, эта функция вычисляет (мультипликативную) обратную величину A (см. torch.linalg.inv()).

Примечание

Рассмотрите использование torch.linalg.tensorsolve(), если это возможно, для умножения тензора слева на обратный тензор, как:

linalg.tensorsolve(A, B) == torch.tensordot(linalg.tensorinv(A), B)  # When B is a tensor with shape A.shape[:B.ndim]

В любом случае предпочтительно использовать tensorsolve(), так как это быстрее и числово устойчивее, чем явное вычисление псевдообратной величины.

См. также

torch.linalg.tensorsolve() вычисляет torch.tensordot(tensorinv(A), B).

Параметры
  • A (Tensor) – тензор для обратного преобразования. Его форма должна удовлетворять условиям prod(A.shape[:ind]) == prod(A.shape[ind:]).
  • ind (int) – индекс, по которому нужно вычислить обратную величину torch.tensordot(). Значение по умолчанию: 2.
Ключевые аргументы

out (Tensor, необязательный) – выходной тензор. Игнорируется, если None. Значение по умолчанию: None.

Возбуждает

RuntimeError – если преобразованный A необратим или произведение первых ind измерений не равно произведению остальных.

Примеры:

>>> A = torch.eye(4 * 6).reshape((4, 6, 8, 3))
>>> Ainv = torch.linalg.tensorinv(A, ind=2)
>>> Ainv.shape
torch.Size([8, 3, 4, 6])
>>> B = torch.randn(4, 6)
>>> torch.allclose(torch.tensordot(Ainv, B), torch.linalg.tensorsolve(A, B))
True

>>> A = torch.randn(4, 4)
>>> Atensorinv = torch.linalg.tensorinv(A, ind=1)
>>> Ainv = torch.linalg.inv(A)
>>> torch.allclose(Atensorinv, Ainv)
True

© 2024, PyTorch Contributors
PyTorch has a BSD-style license, as found in the LICENSE file.
https://pytorch.org/docs/2.1/generated/torch.linalg.tensorinv.html

Spec-Zone.ru

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