diff --git a/solvers/newt_alm.py b/solvers/newt_alm.py index e8103b8..3dafd3b 100644 --- a/solvers/newt_alm.py +++ b/solvers/newt_alm.py @@ -16,6 +16,9 @@ def njit(f): # noqa: F811 return f +MAX_STANDARD_SAMPLES = 10_000 + + class Solver(BaseSolver): name = "Newt-ALM" sampling_strategy = "callback" @@ -34,6 +37,16 @@ def set_objective(self, X, y, alphas, fit_intercept): self.X, self.y, self.lambdas = X, y, alphas self.fit_intercept = fit_intercept + def skip(self, X, y, alphas, fit_intercept): + n_samples = X.shape[0] + if self.inner_solver == "standard" and n_samples > MAX_STANDARD_SAMPLES: + return True, ( + f"{self.name}'s standard inner solver would form a dense " + f"{n_samples}-by-{n_samples} system" + ) + + return False, None + def warm_up(self): self.run_once() @@ -215,7 +228,7 @@ def _compute_direction(self, x, sigma, A, b, y, ATy, lambdas): inner_solver = copy.deepcopy(self.inner_solver) if inner_solver == "auto": - if m > 10000: # Very large m - avoid forming mxm matrices + if m > MAX_STANDARD_SAMPLES: inner_solver = "cg" elif m >= 3 * n and n < 5000: inner_solver = "woodbury" diff --git a/tests/test_newt_alm.py b/tests/test_newt_alm.py new file mode 100644 index 0000000..683ebf2 --- /dev/null +++ b/tests/test_newt_alm.py @@ -0,0 +1,32 @@ +from types import SimpleNamespace + +import pytest + +from solvers.newt_alm import Solver + + +@pytest.mark.parametrize( + ("inner_solver", "n_samples", "expected_skip"), + [ + ("standard", 10_001, True), + ("standard", 10_000, False), + ("auto", 20_000, False), + ], +) +def test_skip_standard_solver_for_large_dense_systems( + inner_solver, + n_samples, + expected_skip, +): + solver = Solver.get_instance(inner_solver=inner_solver) + X = SimpleNamespace(shape=(n_samples, 100)) + + skip, reason = solver.skip( + X=X, + y=None, + alphas=None, + fit_intercept=False, + ) + + assert skip is expected_skip + assert (reason is not None) is expected_skip