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

1"""Primal-dual active-set primitives for the driver loop in :mod:`nncg.solver`. 

2 

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

11 

12from __future__ import annotations 

13 

14import math 

15from collections.abc import Callable 

16from typing import NamedTuple 

17 

18import numpy as np 

19from cvx.linalg import Vector 

20from numpy.typing import NDArray 

21 

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

24 

25ReducedGradient = Callable[["Vector", "Vector | None"], "Vector"] 

26"""Reduced gradient of the subproblem: ``(x, lam) -> s``.""" 

27 

28 

29class _DriveOutcome(NamedTuple): 

30 """Raw outcome of :func:`_drive`, field for field what :class:`nncg.solver.Result` carries. 

31 

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`. 

35 

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

46 

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 

55 

56 

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. 

59 

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] 

67 

68 

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. 

73 

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

82 

83 

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. 

94 

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. 

100 

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 

117 

118 

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. 

123 

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. 

128 

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 

138 

139 

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``. 

148 

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. 

152 

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``. 

159 

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 

173 

174 

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. 

186 

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. 

194 

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. 

205 

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 

215 

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 

229 

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 )