latest (dev)
Copy
Latest development documentation · Updated 2026-10-08
Source code for tensorplay.linalg._solve
"""Inverses, linear solves and least squares."""
import operator
from tensorplay import _C
from tensorplay._C import (
linalg_inv_ex,
linalg_ldl_solve as ldl_solve,
linalg_lu_solve as lu_solve,
linalg_solve_triangular as solve_triangular,
)
from ._common import LstsqResult, SlogdetResult, as_index, check_floating
__all__ = [
"det",
"inv",
"inv_ex",
"ldl_solve",
"lstsq",
"lu_solve",
"pinv",
"slogdet",
"solve",
"solve_ex",
"solve_triangular",
"tensorinv",
"tensorsolve",
]
[docs]
def inv(A):
"""inv(A) -> Tensor"""
return _C.linalg_inv(A)
[docs]
def inv_ex(A, *, check_errors=False):
"""inv_ex(A, *, check_errors=False) -> (Tensor, Tensor)"""
return linalg_inv_ex(A, check_errors=check_errors)
[docs]
def det(A):
"""det(A) -> Tensor"""
return _C.linalg_det(A)
[docs]
def slogdet(A):
"""slogdet(A) -> SlogdetResult(sign, logabsdet)"""
sign, logabsdet = _C.linalg_slogdet(A)
return SlogdetResult(sign, logabsdet)
[docs]
def solve(A, B, *, left=True):
"""solve(A, B, *, left=True) -> Tensor"""
return _C.linalg_solve(A, B, left=left)
[docs]
def solve_ex(A, B, *, left=True, check_errors=False):
"""solve_ex(A, B, *, left=True, check_errors=False) -> (Tensor, Tensor)"""
return _C.linalg_solve_ex(A, B, left=left, check_errors=check_errors)
[docs]
def lstsq(A, B, rcond=None, *, driver=None):
"""lstsq(A, B, rcond=None, *, driver=None) -> LstsqResult(solution, residuals, rank, singular_values)"""
solution, residuals, rank, coefficients = _C.linalg_lstsq(A, B, rcond, driver=driver)
return LstsqResult(solution, residuals, rank, coefficients)
[docs]
def pinv(A, *, atol=None, rtol=None, hermitian=False):
"""pinv(A, *, atol=None, rtol=None, hermitian=False) -> Tensor
Moore-Penrose pseudo-inverse. Singular values (eigenvalue magnitudes
when ``hermitian``) at or below ``max(atol, rtol * sigma_max)`` are
treated as zero; ``rtol`` defaults to ``eps * max(m, n)``, or to zero
when only a positive ``atol`` is given. The tolerances may be floats or
tensors that broadcast against the batch.
"""
check_floating(A, "pinv")
if A.dim() < 2:
raise ValueError("linalg.pinv: input must contain matrices")
if hermitian and A.shape[-1] != A.shape[-2]:
raise ValueError("linalg.pinv: hermitian input must be square")
return _C.linalg_pinv(A, atol=atol, rtol=rtol, hermitian=hermitian)
[docs]
def tensorinv(A, ind=2):
"""tensorinv(A, ind=2) -> Tensor
Inverse of ``A`` seen as a square matrix over the split at ``ind``: the
product of the leading ``ind`` dimensions must equal that of the rest.
"""
ind = as_index(ind, "linalg.tensorinv ind")
if A.dim() < 2:
raise RuntimeError("linalg.tensorinv: input must have at least 2 dimensions")
if ind <= 0 or ind >= A.dim():
raise RuntimeError(
f"linalg.tensorinv: ind must be in [1, {A.dim() - 1}], got {ind}")
shape = list(A.shape)
prod_front = 1
for d in shape[:ind]:
prod_front *= d
prod_tail = 1
for d in shape[ind:]:
prod_tail *= d
if prod_front != prod_tail:
raise RuntimeError(
"linalg.tensorinv: expected an equal product of dimensions on both "
f"sides of ind={ind}, got {prod_front} and {prod_tail}")
Ainv2 = inv(A.reshape(prod_front, prod_tail))
return Ainv2.reshape(shape[ind:] + shape[:ind])
[docs]
def tensorsolve(A, B, dims=None):
"""tensorsolve(A, B, dims=None) -> Tensor
Solves the tensor equation ``A X = B`` after flattening the contracted
dimensions into a square matrix. ``dims`` identifies dimensions of ``A``
that should be moved to the trailing side before the flattening step.
"""
if dims is not None:
try:
moved_dims = [operator.index(dims)]
except TypeError:
try:
moved_dims = [operator.index(d) for d in dims]
except TypeError as exc:
raise TypeError(
"linalg.tensorsolve: dims must contain integers") from exc
ndim = A.dim()
normalized = []
for d in moved_dims:
d = d + ndim if d < 0 else d
if d < 0 or d >= ndim:
raise IndexError(
f"linalg.tensorsolve: dimension {d} out of range for "
f"a {ndim}-D tensor")
if d in normalized:
raise ValueError(
"linalg.tensorsolve: dims must not contain duplicates")
normalized.append(d)
order = [d for d in range(ndim) if d not in normalized]
A = A.permute(order + normalized)
rank_b = B.dim()
if rank_b > A.dim():
raise RuntimeError(
f"linalg.tensorsolve: B with shape {tuple(B.shape)} has more "
f"dimensions than A with shape {tuple(A.shape)}")
if tuple(A.shape[:rank_b]) != tuple(B.shape):
raise RuntimeError(
f"linalg.tensorsolve: B with shape {tuple(B.shape)} must match "
f"the leading dimensions of A {tuple(A.shape[:rank_b])}")
q_shape = list(A.shape[rank_b:])
q_size = 1
for size in q_shape:
q_size *= size
if q_size <= 0 or A.numel() != q_size * q_size or B.numel() != q_size:
raise RuntimeError(
f"linalg.tensorsolve: A with shape {tuple(A.shape)} and B with "
f"shape {tuple(B.shape)} do not form a square tensor equation")
matrix = A.reshape(q_size, q_size)
rhs = B.reshape(q_size)
return solve(matrix, rhs).reshape(q_shape)Help improve this page
Found an error, an unclear step, or a missing example?
Was this page helpful?

