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

1"""Least-squares solver with NaN-aware row filtering.""" 

2 

3from __future__ import annotations 

4 

5import numpy as np 

6import numpy.typing as npt 

7 

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 

12 

13 

14def _condition_number(sv: npt.NDArray[np.floating]) -> float: 

15 """Condition number ``sv[0] / sv[-1]`` from descending singular values. 

16 

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") 

25 

26 

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. 

33 

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. 

39 

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. 

45 

46 Returns: 

47 A four-tuple ``(x, residuals, rank, sv)`` matching the convention of 

48 :func:`numpy.linalg.lstsq`: 

49 

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. 

55 

56 Raises: 

57 DimensionMismatchError: If ``rhs`` length does not match the number of 

58 rows in *matrix*. 

59 

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 

68 

69 NaN rows are silently dropped: 

70 

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]) 

79 

80 n_cols = matrix.shape[1] 

81 

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] 

86 

87 if sub_matrix.shape[0] == 0: 

88 return np.full(n_cols, np.nan), np.array([]), 0, np.array([]) 

89 

90 x, residuals, rank, sv = np.linalg.lstsq(sub_matrix, sub_rhs, rcond=None) 

91 

92 _warn_ill_conditioned(_condition_number(sv), cond_threshold) 

93 

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 )