Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
22 changes: 21 additions & 1 deletion solvers/sortedl1.py
Original file line number Diff line number Diff line change
@@ -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 = {
Expand Down
30 changes: 30 additions & 0 deletions tests/test_sortedl1.py
Original file line number Diff line number Diff line change
@@ -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