diff --git a/docs/source/development/ep-03-alignment.md b/docs/source/development/ep-03-alignment.md index 7c8d03ed6..a1e83f377 100644 --- a/docs/source/development/ep-03-alignment.md +++ b/docs/source/development/ep-03-alignment.md @@ -95,9 +95,11 @@ relevant: - `hess` and `hessp`: Currently we don't support closed form hessians. If we support them they will be called `hess`. In the meantime, this can raise a `NotImplementedError`. -- `callback`: Currently we do not support `callback`s. If we support them they will be - called `callback` and be as compatible with SciPy as possible. In the meantime we can - raise a `NotImplementedError`. +- `callback`: We support SciPy-style callbacks with signature `callback(xk)`, where `xk` + is the current parameter PyTree. In contrast to SciPy, the callback is called after + each objective evaluation and not after each iteration. The + `callback(intermediate_result)` interface and stopping an optimization by raising + `StopIteration` are not yet supported. - If a user sets `jac=True` we raise and error and explain how to use `fun_and_jac` instead. diff --git a/src/optimagic/optimization/create_optimization_problem.py b/src/optimagic/optimization/create_optimization_problem.py index e37d614e5..7fd95c100 100644 --- a/src/optimagic/optimization/create_optimization_problem.py +++ b/src/optimagic/optimization/create_optimization_problem.py @@ -86,6 +86,7 @@ class OptimizationProblem: skip_checks: bool direction: Direction fun_eval: SpecificFunctionValue + callback: Callable[[PyTree], None] | None def create_optimization_problem( @@ -303,14 +304,6 @@ def create_optimization_problem( ) raise NotImplementedError(msg) - if callback is not None: - msg = ( - "The callback argument is not yet supported in optimagic. Creat an issue " - "on https://github.com/optimagic-dev/optimagic/ if you have urgent " - "need for this feature." - ) - raise NotImplementedError(msg) - # ================================================================================== # Handle scipy arguments that will never be supported # ================================================================================== @@ -529,6 +522,21 @@ def create_optimization_problem( if not isinstance(collect_history, bool): raise ValueError("collect_history must be a boolean") + # ================================================================================== + # process and validate callback + # ================================================================================== + + if callback is not None: + if not callable(callback): + raise InvalidFunctionError("callback must be a callable or None.") + # Same signature checks as for fun / jac: one free argument (the params / xk). + callback = partial_func_of_params( + func=callback, + kwargs={}, + name="callback", + skip_checks=skip_checks, + ) + # ================================================================================== # create the problem object # ================================================================================== @@ -551,6 +559,7 @@ def create_optimization_problem( skip_checks=skip_checks, direction=direction, fun_eval=fun_eval, + callback=callback, ) return problem diff --git a/src/optimagic/optimization/internal_optimization_problem.py b/src/optimagic/optimization/internal_optimization_problem.py index 5216bbda5..9c7b52b04 100644 --- a/src/optimagic/optimization/internal_optimization_problem.py +++ b/src/optimagic/optimization/internal_optimization_problem.py @@ -61,6 +61,7 @@ def __init__( linear_constraints: list[dict[str, Any]] | None, nonlinear_constraints: list[dict[str, Any]] | None, logger: LogStore[Any, Any] | None, + callback: Callable[[PyTree], None] | None, # TODO: add hess and hessp ): self._fun = fun @@ -78,6 +79,7 @@ def __init__( self._linear_constraints = linear_constraints self._nonlinear_constraints = nonlinear_constraints self._logger = logger + self._callback = callback self._step_id: int | None = None # ================================================================================== @@ -97,6 +99,7 @@ def fun(self, x: NDArray[np.float64]) -> float | NDArray[np.float64]: """ fun_value, hist_entry = self._evaluate_fun(x) self._history.add_entry(hist_entry) + self._maybe_call_callback(hist_entry.params) return fun_value def jac(self, x: NDArray[np.float64]) -> NDArray[np.float64]: @@ -125,6 +128,7 @@ def fun_and_jac( """ fun_and_jac_value, hist_entry = self._evaluate_fun_and_jac(x) self._history.add_entry(hist_entry) + self._maybe_call_callback(hist_entry.params) return fun_and_jac_value def batch_fun( @@ -158,6 +162,8 @@ def batch_fun( fun_values = [result[0] for result in batch_result] hist_entries = [result[1] for result in batch_result] self._history.add_batch(hist_entries, batch_size) + for hist_entry in hist_entries: + self._maybe_call_callback(hist_entry.params) return fun_values @@ -227,6 +233,8 @@ def batch_fun_and_jac( fun_and_jac_values = [result[0] for result in batch_result] hist_entries = [result[1] for result in batch_result] self._history.add_batch(hist_entries, batch_size) + for hist_entry in hist_entries: + self._maybe_call_callback(hist_entry.params) return fun_and_jac_values @@ -265,6 +273,20 @@ def with_step_id(self, step_id: int) -> Self: new._step_id = step_id return new + def _maybe_call_callback(self, params: PyTree) -> None: + """Call the optional SciPy-style ``callback(xk)`` if one was provided. + + Args: + params: Current external (user-facing) parameters as a PyTree. + + Notes: + Raising ``StopIteration`` from the callback to abort optimization (as in + SciPy) is not handled yet. + + """ + if self._callback is not None: + self._callback(params) + # ================================================================================== # Public attributes # ================================================================================== @@ -958,6 +980,7 @@ def __init__( linear_constraints=linear_constraints, nonlinear_constraints=nonlinear_constraints, logger=logger, + callback=None, ) @@ -1096,4 +1119,5 @@ def derivative_flatten(tree: PyTree, x: NDArray[np.float64]) -> Any: linear_constraints=linear_constraints, nonlinear_constraints=nonlinear_constraints, logger=logger, + callback=None, ) diff --git a/src/optimagic/optimization/optimize.py b/src/optimagic/optimization/optimize.py index 2236269f9..a2c2ecd5d 100644 --- a/src/optimagic/optimization/optimize.py +++ b/src/optimagic/optimization/optimize.py @@ -76,8 +76,7 @@ JacType = Callable[..., PyTree] FunAndJacType = Callable[..., tuple[float | PyTree | FunctionValue, PyTree]] HessType = Callable[..., PyTree] -# TODO: refine this type -CallbackType = Callable[..., Any] +CallbackType = Callable[[PyTree], None] CriterionType = Callable[..., float | dict[str, Any]] CriterionAndDerivativeType = Callable[..., tuple[float | dict[str, Any], PyTree]] @@ -216,7 +215,13 @@ def maximize( args: Alternative to fun_kwargs for scipy compatibility. hess: Not yet supported. hessp: Not yet supported. - callback: Not yet supported. + callback: Optional callable called after each objective evaluation with + signature ``callback(xk)``, where ``xk`` holds the current parameters (a + PyTree with the same structure as ``params``). ``xk`` is not copied, so + the callback must not modify it in place. The callback is not called + during the exploration phase of a multistart optimization. + Raising ``StopIteration`` to abort optimization is not yet supported. The + ``callback(intermediate_result)`` interface is not yet supported. options: Not yet supported. tol: Not yet supported. criterion: Deprecated. Use fun instead. @@ -413,7 +418,13 @@ def minimize( args: Alternative to fun_kwargs for scipy compatibility. hess: Not yet supported. hessp: Not yet supported. - callback: Not yet supported. + callback: Optional callable called after each objective evaluation with + signature ``callback(xk)``, where ``xk`` holds the current parameters (a + PyTree with the same structure as ``params``). ``xk`` is not copied, so + the callback must not modify it in place. The callback is not called + during the exploration phase of a multistart optimization. + Raising ``StopIteration`` to abort optimization is not yet supported. The + ``callback(intermediate_result)`` interface is not yet supported. options: Not yet supported. tol: Not yet supported. criterion: Deprecated. Use fun instead. @@ -653,6 +664,7 @@ def _optimize(problem: OptimizationProblem) -> OptimizeResult: linear_constraints=None, nonlinear_constraints=internal_nonlinear_constraints, logger=logger, + callback=problem.callback, ) # ================================================================================== diff --git a/tests/optimagic/optimization/test_callback.py b/tests/optimagic/optimization/test_callback.py new file mode 100644 index 000000000..e2d2a3d9d --- /dev/null +++ b/tests/optimagic/optimization/test_callback.py @@ -0,0 +1,207 @@ +"""Tests for SciPy-style callback(xk) support.""" + +from dataclasses import dataclass + +import numpy as np +import pytest +from numpy.testing import assert_array_equal as aae + +import optimagic as om +from optimagic import mark +from optimagic.exceptions import InvalidFunctionError, InvalidKwargsError +from optimagic.optimization.algorithm import Algorithm, InternalOptimizeResult +from optimagic.optimization.optimize import minimize +from optimagic.typing import AggregationLevel + + +def test_callback_matches_history_params(): + """Callback parameters match the collected optimization history.""" + xs = [] + + def callback(xk): + xs.append(xk) + + res = om.minimize( + fun=lambda x: x @ x, + x0=np.arange(3, dtype=float), + algorithm="scipy_neldermead", + callback=callback, + ) + + assert res.history is not None + assert len(xs) == len(res.history.params) + aae(xs, res.history.params) + + +def test_callback_receives_external_params(): + """Callback gets external PyTree params, not the internal flat vector.""" + received = [] + + def callback(xk): + received.append(xk) + + params = {"a": np.array([1.0, 2.0]), "b": np.array([3.0])} + + def fun(p): + return p["a"] @ p["a"] + p["b"] @ p["b"] + + res = om.minimize( + fun=fun, + params=params, + algorithm="scipy_neldermead", + callback=callback, + ) + + assert len(received) >= 1 + assert isinstance(received[0], dict) + assert set(received[0]) == {"a", "b"} + assert received[0]["a"].shape == (2,) + assert received[0]["b"].shape == (1,) + + assert res.history is not None + assert len(received) == len(res.history.params) + for got, expected in zip(received, res.history.params, strict=True): + assert set(got) == set(expected) + aae(got["a"], expected["a"]) + aae(got["b"], expected["b"]) + + +@mark.minimizer( + name="dummy_callback_jac", + solver_type=AggregationLevel.SCALAR, + is_available=True, + is_global=False, + needs_jac=True, + needs_hess=False, + needs_bounds=False, + supports_parallelism=False, + supports_bounds=False, + supports_infinite_bounds=False, + supports_linear_constraints=False, + supports_nonlinear_constraints=False, + disable_history=False, +) +@dataclass(frozen=True) +class _DummyJacOptimizer(Algorithm): + def _solve_internal_problem(self, problem, x0): + problem.jac(x0) + problem.jac(x0 + 1) + problem.fun(x0 + 2) + + return InternalOptimizeResult(x=x0 + 2, fun=0.0, success=True) + + +def test_callback_not_called_on_jac(): + """Callback runs on objective evaluations but not on jac-only evaluations.""" + xs = [] + + res = minimize( + fun=lambda x: x @ x, + params=np.arange(3, dtype=float), + algorithm=_DummyJacOptimizer, + callback=xs.append, + ) + + aae(xs, [np.arange(3) + 2.0]) + assert len(res.history.params) == 3 + + +def test_invalid_callback_too_few_arguments(): + msg = "callback must have at least one free argument" + + def bad_callback(): + return None + + with pytest.raises(InvalidFunctionError, match=msg): + om.minimize( + fun=lambda x: x @ x, + x0=np.arange(3, dtype=float), + algorithm="scipy_neldermead", + callback=bad_callback, + ) + + +def test_invalid_callback_too_many_required_arguments(): + msg = "Too few keyword arguments for callback" + + def bad_callback(xk, extra): + return None + + with pytest.raises(InvalidKwargsError, match=msg): + om.minimize( + fun=lambda x: x @ x, + x0=np.arange(3, dtype=float), + algorithm="scipy_neldermead", + callback=bad_callback, + ) + + +def test_invalid_callback_not_callable(): + with pytest.raises(InvalidFunctionError, match="callback must be a callable"): + om.minimize( + fun=lambda x: x @ x, + x0=np.arange(3, dtype=float), + algorithm="scipy_neldermead", + callback="not-a-callable", + ) + + +@mark.minimizer( + name="dummy_callback_parallel", + solver_type=AggregationLevel.SCALAR, + is_available=True, + is_global=False, + needs_jac=False, + needs_hess=False, + needs_bounds=False, + supports_parallelism=True, + supports_bounds=False, + supports_infinite_bounds=False, + supports_linear_constraints=False, + supports_nonlinear_constraints=False, + disable_history=False, +) +@dataclass(frozen=True) +class _DummyParallelOptimizer(Algorithm): + n_cores: int = 1 + batch_size: int = 1 + + def _solve_internal_problem(self, problem, x0): + xs = np.arange(15).repeat(len(x0)).reshape(15, len(x0)) + + for iteration in range(3): + start_index = iteration * 5 + problem.batch_fun( + list(xs[start_index : start_index + 4]), + n_cores=self.n_cores, + batch_size=self.batch_size, + ) + problem.fun(xs[start_index + 4]) + + return InternalOptimizeResult( + x=xs[-1], + fun=5, + success=True, + n_fun_evals=15, + n_iterations=3, + ) + + +def test_callback_history_with_parallel_optimizer(): + """History collected via callback matches optimagic history under parallelism.""" + collected = [] + + def callback(xk): + collected.append(xk) + + res = minimize( + fun=lambda x: 5.0, + params=np.arange(5, dtype=float), + algorithm=_DummyParallelOptimizer, + algo_options={"n_cores": 2, "batch_size": 2}, + callback=callback, + ) + + assert res.history is not None + assert len(collected) == len(res.history.params) + aae(collected, res.history.params) diff --git a/tests/optimagic/optimization/test_internal_optimization_problem.py b/tests/optimagic/optimization/test_internal_optimization_problem.py index 0a8f7bc72..0eeb217b6 100644 --- a/tests/optimagic/optimization/test_internal_optimization_problem.py +++ b/tests/optimagic/optimization/test_internal_optimization_problem.py @@ -74,6 +74,7 @@ def fun_and_jac(params): linear_constraints=linear_constraints, nonlinear_constraints=nonlinear_constraints, logger=None, + callback=None, ) return problem @@ -481,6 +482,7 @@ def derivative_flatten(tree, x): linear_constraints=linear_constraints, nonlinear_constraints=nonlinear_constraints, logger=None, + callback=None, ) return problem @@ -605,6 +607,7 @@ def fun_and_jac(params): linear_constraints=linear_constraints, nonlinear_constraints=nonlinear_constraints, logger=None, + callback=None, ) return problem diff --git a/tests/optimagic/optimization/test_scipy_aliases.py b/tests/optimagic/optimization/test_scipy_aliases.py index 113221674..6720c5ebe 100644 --- a/tests/optimagic/optimization/test_scipy_aliases.py +++ b/tests/optimagic/optimization/test_scipy_aliases.py @@ -120,17 +120,6 @@ def test_exception_for_hessp(): ) -def test_exception_for_callback(): - msg = "The callback argument is not yet supported" - with pytest.raises(NotImplementedError, match=msg): - om.minimize( - fun=lambda x: x @ x, - x0=np.arange(3), - algorithm="scipy_lbfgsb", - callback=print, - ) - - def test_exception_for_options(): msg = "The options argument is not supported" with pytest.raises(NotImplementedError, match=msg):