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.
-
A (Tensor) – тензор для обратного преобразования. Его форма должна удовлетворять условиям
- Ключевые аргументы
-
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