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.inverse(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/1.13/generated/torch.linalg.tensorinv.html