Coverage for src/cvx/linalg/solve/lstsq.py: 100%
25 statements
« prev ^ index » next coverage.py v7.16.2, created at 2026-10-01 07:23 +0000
« prev ^ index » next coverage.py v7.16.2, created at 2026-10-01 07:23 +0000
1"""Least-squares solver with NaN-aware row filtering."""
3from __future__ import annotations
5import numpy as np
6import numpy.typing as npt
8from ..core.condition import DEFAULT_COND_THRESHOLD
9from ..core.condition import warn_ill_conditioned as _warn_ill_conditioned
10from ..core.exceptions import DimensionMismatchError
11from ..core.types import Matrix, Vector
14def _condition_number(sv: npt.NDArray[np.floating]) -> float:
15 """Condition number ``sv[0] / sv[-1]`` from descending singular values.
17 Returns ``inf`` when the smallest singular value is zero and ``1.0`` when
18 there are no singular values (an empty valid sub-matrix).
19 """
20 if sv.size == 0:
21 return 1.0
22 if sv[-1] > 0:
23 return float(sv[0] / sv[-1])
24 return float("inf")
27def lstsq(
28 matrix: Matrix,
29 rhs: Vector,
30 cond_threshold: float | None = DEFAULT_COND_THRESHOLD,
31) -> tuple[Vector, Vector, int, Vector]:
32 """Solve an overdetermined or underdetermined system in the least-squares sense.
34 Rows where any entry in *matrix* or the corresponding entry in *rhs* is
35 non-finite are excluded before solving. The returned solution vector
36 always has length equal to the number of columns in *matrix*. When the
37 effective condition number of the valid sub-matrix exceeds
38 *cond_threshold*, an ``IllConditionedMatrixWarning`` is emitted.
40 Args:
41 matrix: Coefficient matrix of shape ``(m, n)``.
42 rhs: Right-hand side vector of length ``m``.
43 cond_threshold: Condition-number threshold above which a warning is
44 emitted. Defaults to ``1e12``. ``None`` skips the check.
46 Returns:
47 A four-tuple ``(x, residuals, rank, sv)`` matching the convention of
48 :func:`numpy.linalg.lstsq`:
50 - ``x`` — least-squares solution of shape ``(n,)``.
51 - ``residuals`` — sum of squared residuals; empty when the solution is
52 not unique or all rows are invalid.
53 - ``rank`` — effective rank of the valid sub-matrix.
54 - ``sv`` — singular values of the valid sub-matrix in descending order.
56 Raises:
57 DimensionMismatchError: If ``rhs`` length does not match the number of
58 rows in *matrix*.
60 Example:
61 >>> import numpy as np
62 >>> from cvx.linalg import lstsq
63 >>> A = np.array([[1.0, 1.0], [1.0, 2.0], [1.0, 3.0]])
64 >>> b = np.array([6.0, 5.0, 7.0])
65 >>> x, res, rank, sv = lstsq(A, b)
66 >>> int(rank)
67 2
69 NaN rows are silently dropped:
71 >>> A_nan = np.array([[1.0, 1.0], [np.nan, 2.0], [1.0, 3.0]])
72 >>> b_nan = np.array([6.0, 5.0, 7.0])
73 >>> x2, _, rank2, _ = lstsq(A_nan, b_nan)
74 >>> int(rank2)
75 2
76 """
77 if rhs.shape[0] != matrix.shape[0]:
78 raise DimensionMismatchError(rhs.shape[0], matrix.shape[0])
80 n_cols = matrix.shape[1]
82 # Filter rows that contain any non-finite value in matrix or rhs.
83 row_mask = np.isfinite(matrix).all(axis=1) & np.isfinite(rhs)
84 sub_matrix = matrix[row_mask]
85 sub_rhs = rhs[row_mask]
87 if sub_matrix.shape[0] == 0:
88 return np.full(n_cols, np.nan), np.array([]), 0, np.array([])
90 x, residuals, rank, sv = np.linalg.lstsq(sub_matrix, sub_rhs, rcond=None)
92 _warn_ill_conditioned(_condition_number(sv), cond_threshold)
94 return (
95 x.astype(np.float64, copy=False),
96 residuals.astype(np.float64, copy=False),
97 int(rank),
98 sv.astype(np.float64, copy=False),
99 )