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
« 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.
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).
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
14 min 1/2 w^T Sigma w - lam mu^T w s.t. ||w||_1 <= c, A w = 0
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.
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"""
31from __future__ import annotations
33import itertools
35import numpy as np
36from numpy.typing import NDArray
37from scipy.optimize import linprog # type: ignore[import-untyped]
39from .cla import CLA
40from .operators import QuadraticForm
41from .operators._core import _RCOND_FLOOR
43#: One LASSO breakpoint as ``(lam, beta, active)``; :mod:`cvxcla.lasso` wraps it.
44BreakpointData = tuple[float, NDArray[np.float64], NDArray[np.bool_]]
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``).
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)
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``.
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``).
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)
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.
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.
109 Returns:
110 The breakpoints ``(lam, beta, active)`` from ``lam_max`` down to the
111 constrained least-squares fit at ``lam = 0``.
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))]
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
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