From e95ed83f6ae7d41105cc778f9e3c15a524f7e8bf Mon Sep 17 00:00:00 2001 From: Johan Larsson Date: Sun, 30 Aug 2026 12:46:09 +0200 Subject: [PATCH] refactor: use min tolerance for sortedl1 --- solvers/sortedl1.py | 22 +++++++++++++++++++++- tests/test_sortedl1.py | 30 ++++++++++++++++++++++++++++++ 2 files changed, 51 insertions(+), 1 deletion(-) create mode 100644 tests/test_sortedl1.py diff --git a/solvers/sortedl1.py b/solvers/sortedl1.py index ed5f326..189e480 100644 --- a/solvers/sortedl1.py +++ b/solvers/sortedl1.py @@ -1,14 +1,34 @@ from benchopt import BaseSolver, safe_import_context -from benchopt.stopping_criterion import INFINITY +from benchopt.stopping_criterion import INFINITY, SufficientProgressCriterion with safe_import_context() as import_ctx: import numpy as np from sortedl1 import Slope +MIN_TOLERANCE = 1e-6 + + +class MinimumToleranceCriterion(SufficientProgressCriterion): + def __init__(self, min_tol=MIN_TOLERANCE, **kwargs): + super().__init__(**kwargs) + self.min_tol = min_tol + self.kwargs["min_tol"] = min_tol + + def should_stop(self, stop_val, objective_list): + stop, status, next_stop_val = super().should_stop( + stop_val, + objective_list, + ) + if not stop and next_stop_val < self.min_tol: + return True, "done", stop_val + return stop, status, next_stop_val + + class Solver(BaseSolver): name = "sortedl1" sampling_strategy = "tolerance" + stopping_criterion = MinimumToleranceCriterion() install_cmd = "conda" requirements = ["pip::sortedl1"] parameters = { diff --git a/tests/test_sortedl1.py b/tests/test_sortedl1.py new file mode 100644 index 0000000..c98733c --- /dev/null +++ b/tests/test_sortedl1.py @@ -0,0 +1,30 @@ +import pytest + +from solvers.sortedl1 import MIN_TOLERANCE, Solver + + +def test_tolerance_sampling_stops_at_minimum_tolerance(): + solver = Solver.get_instance() + criterion = solver.stopping_criterion.get_runner_instance( + solver=solver, + max_runs=100, + ) + stop_val = criterion.init_stop_val() + evaluated_tolerances = [] + objective_values = [] + + while True: + evaluated_tolerances.append(stop_val) + objective_values.append( + {"objective_value": -float(len(evaluated_tolerances))} + ) + stop, status, stop_val = criterion.should_stop( + stop_val, + objective_values, + ) + if stop: + break + + assert status == "done" + assert evaluated_tolerances[-1] == pytest.approx(MIN_TOLERANCE) + assert min(evaluated_tolerances) >= MIN_TOLERANCE