Coverage for src/cvxcla/_lasso_cla.py: 100%

40 statements  

« prev     ^ index     » next       coverage.py v7.16.2, created at 2026-09-29 05:24 +0000

1"""The equality-constrained LASSO ``A beta = 0``, traced through the leverage CLA. 

2 

3The LASSO homotopy of :mod:`cvxcla.lasso` starts from ``beta = 0`` and needs every 

4constraint row slack there, so it cannot carry equality rows: an equality row is 

5active from the start and the path cannot be seeded one coordinate at a time. The 

6Critical Line Algorithm has no such restriction, and under ``Sigma = X^T X`` and 

7``mu = X^T y`` it traces the same curve (Schmelzer and Hastie, "The Critical Line 

8Algorithm and the Constrained LASSO: One Curve, Two Literatures", 

9arXiv:2609.25704). 

10 

11With homogeneous constraints the route is Corollary 2 of that note. At a fixed 

12gross-exposure cap ``c`` the tilt sweep of the capped program 

13 

14 min 1/2 w^T Sigma w - lam mu^T w s.t. ||w||_1 <= c, A w = 0 

15 

16is the budget-indexed path rescaled, ``w(lam) = lam * beta`` with 

17``||beta||_1 = c / lam``, and by Theorem 1 the budget-indexed path is the 

18constrained LASSO path. So one leverage-capped CLA trace with ``c = 1`` gives 

19every LASSO breakpoint as ``beta = w / lam``, in order: ``lam = inf`` is 

20``beta = 0``, and the cap's release is the constrained least-squares fit. A box of 

21``+/- 2 c`` makes the CLA well posed and never binds, since ``|w_i| <= ||w||_1 <= c``. 

22Under ``beta >= 0`` the lower bound is ``0``, which is homogeneous as well. 

23 

24The CLA's tilt is not the LASSO penalty. The penalty at a breakpoint is read off 

25the KKT conditions of the LASSO on the segment that leaves it: 

26``(X^T y - X^T X beta)_F = lam_L s_F + A_F^T nu``, over the segment's free set 

27``F`` with signs ``s_F``. The matrix ``[s_F, A_F^T]`` has full column rank at 

28every non-degenerate active set, so the least-squares solution is exact. 

29""" 

30 

31from __future__ import annotations 

32 

33import itertools 

34 

35import numpy as np 

36from numpy.typing import NDArray 

37from scipy.optimize import linprog # type: ignore[import-untyped] 

38 

39from .cla import CLA 

40from .operators import QuadraticForm 

41from .operators._core import _RCOND_FLOOR 

42 

43#: One LASSO breakpoint as ``(lam, beta, active)``; :mod:`cvxcla.lasso` wraps it. 

44BreakpointData = tuple[float, NDArray[np.float64], NDArray[np.bool_]] 

45 

46 

47def _max_correlation(xty: NDArray[np.float64], a: NDArray[np.float64], nonneg: bool) -> float: 

48 """Return ``max xty^T beta`` over ``||beta||_1 <= 1``, ``A beta = 0`` (and ``beta >= 0``). 

49 

50 Zero means ``beta = 0`` is optimal at every penalty, so the path is one point. 

51 """ 

52 n = xty.shape[0] 

53 legs = np.hstack([np.eye(n), -np.eye(n)])[:, : n if nonneg else 2 * n] 

54 result = linprog( 

55 c=-(xty @ legs), 

56 A_ub=np.ones((1, legs.shape[1])), 

57 b_ub=np.ones(1), 

58 A_eq=a @ legs, 

59 b_eq=np.zeros(a.shape[0]), 

60 bounds=(0.0, None), 

61 method="highs", 

62 ) 

63 return float(-result.fun) 

64 

65 

66def _penalty( 

67 quad: QuadraticForm, 

68 xty: NDArray[np.float64], 

69 a: NDArray[np.float64], 

70 beta: NDArray[np.float64], 

71 free: NDArray[np.bool_], 

72 signs: NDArray[np.float64], 

73) -> float: 

74 """Solve the LASSO stationarity on ``free`` for the penalty ``lam_L``. 

75 

76 Args: 

77 quad: The Gram form ``X^T X``. 

78 xty: The linear term ``X^T y``. 

79 a: The equality rows. 

80 beta: The coefficients at the breakpoint. 

81 free: The support of the segment leaving the breakpoint. 

82 signs: The signs of ``beta`` on that segment (length ``n``). 

83 

84 Returns: 

85 The penalty, clamped at zero against round-off at the least-squares end. 

86 """ 

87 residual = (xty - quad.matvec(beta))[free] 

88 system = np.hstack([signs[free][:, None], a[:, free].T]) 

89 solution = np.linalg.lstsq(system, residual, rcond=None)[0] 

90 return max(float(solution[0]), 0.0) 

91 

92 

93def equality_path( 

94 quad: QuadraticForm, 

95 xty: NDArray[np.float64], 

96 a: NDArray[np.float64], 

97 nonneg: bool, 

98 tol: float, 

99) -> list[BreakpointData]: 

100 """Trace the LASSO path under ``A beta = 0`` via one leverage-capped CLA. 

101 

102 Args: 

103 quad: The Gram form ``X^T X`` as a :class:`QuadraticForm`. 

104 xty: The linear term ``X^T y``. 

105 a: The equality rows ``(m, n)`` of ``A beta = 0``. 

106 nonneg: Restrict to ``beta >= 0``. 

107 tol: Below this, the largest feasible correlation counts as zero. 

108 

109 Returns: 

110 The breakpoints ``(lam, beta, active)`` from ``lam_max`` down to the 

111 constrained least-squares fit at ``lam = 0``. 

112 

113 Raises: 

114 ValueError: If the Gram form is singular, so the constrained least-squares 

115 fit is not unique (e.g. more features than observations), or if the CLA 

116 cannot trace the problem (e.g. a degenerate first vertex). 

117 """ 

118 n = xty.shape[0] 

119 if quad.rcond_free(np.arange(n)) < _RCOND_FLOOR: 

120 msg = ( 

121 "the equality-constrained LASSO needs a positive-definite Gram X^T X " 

122 "(more observations than features in general position), so that the " 

123 "constrained least-squares end of the path is unique" 

124 ) 

125 raise ValueError(msg) 

126 if _max_correlation(xty, a, nonneg) <= tol: 

127 return [(0.0, np.zeros(n), np.zeros(n, dtype=bool))] 

128 

129 try: 

130 cla = CLA( 

131 mean=xty, 

132 covariance=quad, 

133 lower_bounds=np.zeros(n) if nonneg else np.full(n, -2.0), 

134 upper_bounds=np.full(n, 2.0), 

135 a=a, 

136 b=np.zeros(a.shape[0]), 

137 leverage=1.0, 

138 ) 

139 except ValueError as err: 

140 msg = f"the equality-constrained LASSO could not be traced through the leverage CLA: {err}" 

141 raise ValueError(msg) from err 

142 

143 tps = cla.turning_points 

144 path: list[BreakpointData] = [] 

145 # Segment k runs from turning point k to k + 1; the last one ends at the CLA's 

146 # lambda = 0 endpoint w = 0, which has no beta of its own and is not recorded. 

147 for hi, lo in itertools.pairwise(tps): 

148 if np.isinf(hi.lamb): 

149 # The first segment holds the maximum-return vertex, so beta runs along 

150 # the ray from 0 through it and carries the vertex's signs. 

151 beta, signs = np.zeros(n), np.sign(hi.weights) 

152 else: 

153 # w is affine in lambda on the segment, so this is beta at its midpoint, 

154 # where every free coordinate is strictly nonzero. 

155 beta = hi.weights / hi.lamb 

156 signs = np.sign((hi.weights + lo.weights) / (hi.lamb + lo.lamb)) 

157 path.append((_penalty(quad, xty, a, beta, hi.free, signs), beta, hi.free.copy())) 

158 return path