Coverage for src/nncg/_equality.py: 100%

29 statements  

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

1"""Equality-augmented reduction: the Schur-complement saddle solve on a free set. 

2 

3Factored out of :meth:`nncg.solver.ActiveSetSolver.solve_eq` so the outer loop 

4keeps only its orchestration. On a free set the saddle system for ``B x = c`` is 

5solved by eliminating the multiplier ``lambda`` in R^p through the p-by-p Schur 

6complement ``S = B_F A_F^{-1} B_F^T``: the ``p + 1`` right-hand sides share the 

7operator ``A_F`` and are each one inner solve, then ``S lambda = c - B_F v0`` 

8fixes the multipliers in closed form. 

9 

10Also home to the equality-specific input check (:func:`_require_eq_shapes`) and 

11the feasibility certificate (:func:`_eq_feasible`) that ``solve_eq`` adds on top 

12of the KKT exit: a rank-deficient, inconsistent ``B`` can pass the KKT test 

13while ``B x != c``, and must not be reported as converged. 

14""" 

15 

16from __future__ import annotations 

17 

18from typing import TYPE_CHECKING 

19 

20import numpy as np 

21from cvx.linalg import Matrix, SymmetricOperator, Vector, cholesky_solve 

22from numpy.typing import NDArray 

23 

24if TYPE_CHECKING: 

25 from .solver import InnerSolver 

26 

27 

28def _saddle_solve( 

29 inner: InnerSolver, 

30 a: SymmetricOperator, 

31 b: Vector, 

32 b_eq: Matrix, 

33 c_eq: Vector, 

34 idx: NDArray[np.int_], 

35 x0: Vector | None, 

36) -> tuple[Vector, Vector, int]: 

37 """Solve the equality-augmented saddle system on the free set ``idx``. 

38 

39 Runs ``p + 1`` inner solves through the shared free-block operator ``A_F`` 

40 (the ``v0`` column warm-started at ``x0``, the ``v1`` columns cold), forms the 

41 SPD Schur complement ``S = B_F A_F^{-1} B_F^T`` and recovers the multipliers 

42 from ``S lambda = c - B_F v0`` before back-substituting ``x_F = v0 + v1 lambda``. 

43 

44 Args: 

45 inner: The inner solver driving each free-block solve. 

46 a: The SPD operator ``A``. 

47 b: The linear term ``b``. 

48 b_eq: Equality matrix ``B`` of shape ``(p, n)``, full row rank on ``idx``. 

49 c_eq: Equality right-hand side ``c`` of shape ``(p,)``. 

50 idx: Integer positions of the free set ``F``. 

51 x0: Warm inner guess for the ``v0`` column restricted to ``idx``, or ``None``. 

52 

53 Returns: 

54 ``(x_F, lam, inner_iters)``: the free-block solution, the equality 

55 multipliers, and the total inner iteration count across all columns. 

56 """ 

57 p = b_eq.shape[0] 

58 b_f = b_eq[:, idx] 

59 v0, k0 = inner.solve(a, idx, b[idx], x0) 

60 v1 = np.zeros((idx.size, p)) 

61 k_cols = 0 

62 for j in range(p): 

63 v1[:, j], kj = inner.solve(a, idx, b_f[j], None) 

64 k_cols += kj 

65 schur = b_f @ v1 # p-by-p Schur complement, SPD 

66 lam = cholesky_solve(schur, c_eq - b_f @ v0) 

67 xf = v0 + v1 @ lam # x_F = A_F^{-1}(b_F + B_F^T lambda) 

68 return xf, lam, k0 + k_cols 

69 

70 

71def _require_eq_shapes(b_eq: Matrix, c_eq: Vector, n: int) -> None: 

72 """Validate that ``B`` has shape ``(p, n)`` and ``c`` has shape ``(p,)``. 

73 

74 Args: 

75 b_eq: Equality matrix ``B``. 

76 c_eq: Equality right-hand side ``c``. 

77 n: Problem dimension ``len(b)``. 

78 

79 Raises: 

80 ValueError: When ``b_eq`` is not two-dimensional with ``n`` columns, or 

81 ``c_eq`` is not one-dimensional with one entry per row of ``b_eq``. 

82 """ 

83 b_shape, c_shape = np.shape(b_eq), np.shape(c_eq) 

84 if len(b_shape) != 2 or b_shape[1] != n: 

85 msg = f"b_eq must have shape (p, {n}), got {b_shape}" 

86 raise ValueError(msg) 

87 if c_shape != (b_shape[0],): 

88 msg = f"c_eq must have shape ({b_shape[0]},) to match b_eq {b_shape}, got {c_shape}" 

89 raise ValueError(msg) 

90 

91 

92def _eq_feasible(b_eq: Matrix, c_eq: Vector, x: Vector, tol: float) -> bool: 

93 """Return whether ``x`` satisfies ``B x = c`` to the relative tolerance ``tol``. 

94 

95 With ``B_F`` of full row rank the Schur-complement solve makes ``B x = c`` 

96 hold to rounding, so this only fails when the rank precondition is violated 

97 and ``c`` lies outside the range of ``B_F`` — the case the KKT test alone 

98 would certify. 

99 

100 Args: 

101 b_eq: Equality matrix ``B`` of shape ``(p, n)``. 

102 c_eq: Equality right-hand side ``c`` of shape ``(p,)``. 

103 x: The candidate solution. 

104 tol: Tolerance on ``max|B x - c|``, relative to ``max(1, max|c|)``. 

105 

106 Returns: 

107 True when the equality residual is within tolerance. 

108 """ 

109 scale = max(1.0, float(np.max(np.abs(c_eq), initial=0.0))) 

110 return float(np.max(np.abs(b_eq @ x - c_eq), initial=0.0)) <= tol * scale