Coverage for src/nncg/_active_set.py: 100%
75 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"""Primal-dual active-set primitives for the driver loop in :mod:`nncg.solver`.
3The pure working-set algebra the outer loop is built from, factored out of the
4:class:`~nncg.solver.ActiveSetSolver` orchestration: seeding the free set from a
5cold or warm start (:func:`_init_active_set`), splitting the KKT violators into
6primal and dual sets (:func:`_violators`), applying one guarded working-set pivot
7(:func:`_pivot`), and the per-outer-step scaffolding (:func:`_init_run_state`,
8:func:`_solve_free_set`). None of it knows about preconditioning or the inner
9solver — it operates on boolean free masks and index arrays alone.
10"""
12from __future__ import annotations
14import math
15from collections.abc import Callable
16from typing import NamedTuple
18import numpy as np
19from cvx.linalg import Vector
20from numpy.typing import NDArray
22SubSolve = Callable[[NDArray[np.int_], "Vector | None"], "tuple[Vector, Vector | None, int]"]
23"""Subproblem solve on a free set: ``(idx, x0) -> (x_F, lam, inner_iters)``."""
25ReducedGradient = Callable[["Vector", "Vector | None"], "Vector"]
26"""Reduced gradient of the subproblem: ``(x, lam) -> s``."""
29class _DriveOutcome(NamedTuple):
30 """Raw outcome of :func:`_drive`, field for field what :class:`nncg.solver.Result` carries.
32 Named rather than positional so the caller cannot silently swap the three
33 adjacent ``int`` counters. Kept here, not in :mod:`nncg.solver`, so this
34 module need not depend on :class:`~nncg.solver.Result`.
36 Attributes:
37 x: Final iterate.
38 outer: Outer (active-set) steps taken.
39 inner: Total inner iterations across all subproblem solves.
40 fallback: Number of least-index Bland fallback pivots fired.
41 converged: Whether KKT was certified before the outer cap.
42 free: Final free mask.
43 lam: Equality multipliers from the last subproblem (None if bound-only).
44 traj: Visited free-set trajectory when tracked, else None.
45 """
47 x: Vector
48 outer: int
49 inner: int
50 fallback: int
51 converged: bool
52 free: NDArray[np.bool_]
53 lam: Vector | None
54 traj: list[tuple[int, ...]] | None
57def _init_active_set(n: int, warm: tuple[NDArray[np.bool_], Vector] | None) -> tuple[NDArray[np.bool_], Vector | None]:
58 """Seed the free set and inner warm guess for the outer loop.
60 Cold (``warm is None``): everything free (``F = {1..n}``) with no inner
61 guess. Warm: a copy of the previous free mask and the previous iterate as
62 the seed for every subproblem solve.
63 """
64 if warm is None:
65 return np.ones(n, dtype=bool), None # F = {1..n} initially
66 return warm[0].copy(), warm[1]
69def _violators(
70 free: NDArray[np.bool_], x: Vector, s: Vector, tol: float
71) -> tuple[NDArray[np.int_], NDArray[np.int_], NDArray[np.int_]]:
72 """Split the KKT violators at tolerance ``tol`` into primal and dual sets.
74 Returns ``(prim, dual, viol)``: ``prim`` (set ``D``) are free indices whose
75 primal value went negative, ``dual`` (set ``V``) are bound indices whose
76 reduced gradient went negative, and ``viol`` is their concatenation. An empty
77 ``viol`` certifies the KKT conditions at the unique global minimiser.
78 """
79 prim = np.flatnonzero(free & (x < -tol)) # D: free but negative
80 dual = np.flatnonzero((~free) & (s < -tol)) # V: bound but s < 0
81 return prim, dual, np.concatenate([prim, dual])
84def _pivot(
85 free: NDArray[np.bool_],
86 prim: NDArray[np.int_],
87 dual: NDArray[np.int_],
88 viol: NDArray[np.int_],
89 n_bar: int,
90 patience: int,
91 p_max: int,
92) -> tuple[int, int, int]:
93 """Apply one working-set pivot, mutating ``free`` in place.
95 Takes the fast batch exchange — drop every primal violator ``D``, add every
96 dual violator ``V`` — while the violator count strictly drops below ``n_bar``
97 (patience reset to ``p_max``) or patience remains (decremented). Once patience
98 is exhausted without progress it falls back to a single least-index Bland
99 pivot, the load-bearing anti-cycling guarantee behind finite termination.
101 Returns the updated ``(n_bar, patience, fallback_increment)``; the last is 1
102 when the Bland fallback fired, 0 on the batch fast path.
103 """
104 n_viol = viol.size
105 if n_viol < n_bar or patience > 0: # fast path: progress, or patience remains
106 if n_viol < n_bar:
107 n_bar = n_viol
108 patience = p_max
109 else:
110 patience -= 1
111 free[prim] = False # batch exchange: drop all D, add all V
112 free[dual] = True
113 return n_bar, patience, 0
114 i_star = int(np.min(viol)) # anti-cycling fallback: single Bland least-index pivot
115 free[i_star] = not free[i_star]
116 return n_bar, patience, 1
119def _init_run_state(
120 n: int, track: bool, max_outer: int | None, warm: tuple[NDArray[np.bool_], Vector] | None
121) -> tuple[NDArray[np.bool_], Vector | None, list[tuple[int, ...]] | None, float]:
122 """Seed the free set, inner guess, trajectory log and outer-iteration cap.
124 Delegates the free set and warm inner guess to :func:`_init_active_set`,
125 allocates the trajectory list only when ``track`` is set, and resolves
126 ``max_outer`` (``None`` for uncapped) into a numeric loop bound so the
127 driver's ``while`` condition stays branch-free.
129 Returns:
130 ``(free, x_guess, traj, cap)``: the initial free mask, the inner warm
131 guess (``None`` on a cold start), the trajectory list (or ``None``),
132 and the outer-iteration cap (``math.inf`` when uncapped).
133 """
134 free, x_guess = _init_active_set(n, warm)
135 traj: list[tuple[int, ...]] | None = [] if track else None
136 cap = math.inf if max_outer is None else max_outer
137 return free, x_guess, traj, cap
140def _solve_free_set(
141 free: NDArray[np.bool_],
142 n: int,
143 sub_solve: SubSolve,
144 x_guess: Vector | None,
145 traj: list[tuple[int, ...]] | None,
146) -> tuple[Vector, Vector | None, int]:
147 """Solve the subproblem on the current free set and scatter it into ``R^n``.
149 Records the free set on ``traj`` when tracking, seeds the inner solve
150 from ``x_guess`` restricted to the free set (cold when ``None``), and
151 places the returned free-block solution back into a full zero vector.
153 Args:
154 free: Boolean free-set mask over the ``n`` variables.
155 n: Problem dimension.
156 sub_solve: Subproblem callback ``(idx, x0) -> (x_F, lam, inner_iters)``.
157 x_guess: Inner warm guess over all variables, or ``None`` (cold).
158 traj: Trajectory list to append the free set to, or ``None``.
160 Returns:
161 ``(x, lam, inner_iters)``: the full-length iterate, the equality
162 multipliers (``None`` for the bound-only problem), and the inner
163 iteration count from the callback.
164 """
165 idx = np.flatnonzero(free)
166 if traj is not None:
167 traj.append(tuple(idx.tolist()))
168 x0 = x_guess[idx] if x_guess is not None else None
169 xf, lam, k_step = sub_solve(idx, x0)
170 x: Vector = np.zeros(n)
171 x[idx] = xf
172 return x, lam, k_step
175def _drive(
176 tol: float,
177 p_max: int,
178 track: bool,
179 max_outer: int | None,
180 n: int,
181 sub_solve: SubSolve,
182 reduced_gradient: ReducedGradient,
183 warm: tuple[NDArray[np.bool_], Vector] | None,
184) -> _DriveOutcome:
185 """Run the guarded primal-dual active-set loop and return its raw outcome.
187 Owns everything the termination proof depends on: the primal and dual
188 violator tests (:func:`_violators`), the batch exchange with its patience
189 counter and the least-index Bland fallback (:func:`_pivot`). What is solved
190 on each free set enters through ``sub_solve``, with ``reduced_gradient``
191 supplying the matching dual test quantity; the thresholds come straight from
192 :class:`nncg.solver.ActiveSetConfig`. Returns a :class:`_DriveOutcome` so this
193 module need not depend on :class:`nncg.solver.Result` — the caller wraps it.
195 Args:
196 tol: Violator tolerance of the primal and dual KKT tests.
197 p_max: Patience budget before the least-index Bland fallback pivot.
198 track: Record the visited free-set trajectory.
199 max_outer: Optional cap on outer steps (``None`` for uncapped).
200 n: Problem dimension.
201 sub_solve: Callback ``(idx, x0) -> (x_F, lam, inner_iters)`` solving the
202 subproblem on the free set ``idx``.
203 reduced_gradient: Callback ``(x, lam) -> s`` for the dual violator test.
204 warm: Optional ``(free_mask, x_prev)`` pair from a previous solve.
206 Returns:
207 A :class:`_DriveOutcome` with the final iterate, counters and free set.
208 """
209 free, x_guess, traj, cap = _init_run_state(n, track, max_outer, warm)
210 x: Vector = np.zeros(n)
211 lam: Vector | None = None
212 n_bar, patience = n + 1, p_max
213 outer = inner_total = fallback = 0
214 converged = True
216 while outer < cap:
217 x, lam, k_step = _solve_free_set(free, n, sub_solve, x_guess, traj)
218 outer += 1
219 inner_total += k_step
220 if x_guess is not None:
221 x_guess = x # warm mode: newest iterate seeds the next reduced solve
222 prim, dual, viol = _violators(free, x, reduced_gradient(x, lam), tol)
223 if viol.size == 0:
224 break # KKT satisfied -> unique global minimiser
225 n_bar, patience, fired = _pivot(free, prim, dual, viol, n_bar, patience, p_max)
226 fallback += fired
227 else:
228 converged = False # outer cap reached without certifying KKT
230 return _DriveOutcome(
231 x=x,
232 outer=outer,
233 inner=inner_total,
234 fallback=fallback,
235 converged=converged,
236 free=free,
237 lam=lam,
238 traj=traj,
239 )