Coverage for src/cvx/linalg/solve/inv.py: 100%

22 statements  

« prev     ^ index     » next       coverage.py v7.16.2, created at 2026-10-01 07:23 +0000

1"""Matrix inversion helpers that ignore rows and columns with non-finite diagonals.""" 

2 

3from __future__ import annotations 

4 

5import numpy as np 

6 

7from ..core.condition import DEFAULT_COND_THRESHOLD 

8from ..core.condition import ( 

9 check_and_warn_condition as _check_and_warn_condition, 

10) 

11from ..core.exceptions import ( 

12 NonSquareMatrixError, 

13 SingularMatrixError, 

14) 

15from ..core.types import Matrix 

16from ..core.valid import valid 

17 

18 

19def inv( 

20 matrix: Matrix, 

21 cond_threshold: float | None = DEFAULT_COND_THRESHOLD, 

22) -> Matrix: 

23 """Invert a matrix restricted to the valid submatrix. 

24 

25 Rows and columns with non-finite diagonal entries are excluded from the 

26 inversion; the corresponding rows and columns in the result are set to NaN. 

27 When the condition number of the valid sub-matrix exceeds *cond_threshold*, 

28 an ``IllConditionedMatrixWarning`` is emitted. 

29 

30 Args: 

31 matrix: Square matrix to invert. 

32 cond_threshold: Condition-number threshold above which a warning is 

33 emitted. Defaults to ``1e12``. ``None`` skips the check, and the 

34 SVD that computes the condition number, entirely. 

35 

36 Returns: 

37 An inverted matrix with the same shape as *matrix*. Rows and columns 

38 mapped to invalid entries are returned as ``NaN``. 

39 

40 Raises: 

41 NonSquareMatrixError: If the matrix is not square. 

42 SingularMatrixError: If the valid sub-matrix is singular. 

43 

44 Example: 

45 >>> import numpy as np 

46 >>> from cvx.linalg import inv 

47 >>> np.allclose(inv(np.eye(2)), np.eye(2)) 

48 True 

49 

50 NaN-masked entries are skipped: 

51 

52 >>> matrix = np.array([[4.0, 0.0], [0.0, np.nan]]) 

53 >>> result = inv(matrix) 

54 >>> float(result[0, 0]) 

55 0.25 

56 >>> bool(np.isnan(result[0, 1]) and np.isnan(result[1, 0]) and np.isnan(result[1, 1])) 

57 True 

58 """ 

59 if matrix.shape[0] != matrix.shape[1]: 

60 raise NonSquareMatrixError(matrix.shape[0], matrix.shape[1]) 

61 

62 n = matrix.shape[0] 

63 result = np.full((n, n), np.nan) 

64 mask, submatrix = valid(matrix) 

65 

66 if mask.any(): 

67 _check_and_warn_condition(submatrix, cond_threshold) 

68 try: 

69 sub_inv = np.linalg.inv(submatrix) 

70 except np.linalg.LinAlgError as exc: 

71 raise SingularMatrixError(str(exc)) from exc 

72 

73 idx = np.where(mask)[0] 

74 result[np.ix_(idx, idx)] = sub_inv 

75 

76 return result