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
« 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.
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.
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"""
16from __future__ import annotations
18from typing import TYPE_CHECKING
20import numpy as np
21from cvx.linalg import Matrix, SymmetricOperator, Vector, cholesky_solve
22from numpy.typing import NDArray
24if TYPE_CHECKING:
25 from .solver import InnerSolver
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``.
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``.
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``.
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
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,)``.
74 Args:
75 b_eq: Equality matrix ``B``.
76 c_eq: Equality right-hand side ``c``.
77 n: Problem dimension ``len(b)``.
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)
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``.
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.
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|)``.
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