From eb11c0127ef26f211bb723cc6e8a614517dcac6f Mon Sep 17 00:00:00 2001 From: Pragati Date: Sun, 2 Aug 2026 22:33:35 +0530 Subject: [PATCH 1/5] ENH: support SciPy-style callback(xk) during optimization Pass callback through to InternalOptimizationProblem and invoke it on objective evaluations so it works for all optimizers. --- .../create_optimization_problem.py | 10 ++----- .../internal_optimization_problem.py | 10 +++++++ src/optimagic/optimization/optimize.py | 11 ++++++-- .../optimization/test_scipy_aliases.py | 26 ++++++++++++------- 4 files changed, 38 insertions(+), 19 deletions(-) diff --git a/src/optimagic/optimization/create_optimization_problem.py b/src/optimagic/optimization/create_optimization_problem.py index e37d614e5..51eca0baa 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[[Any], Any] | 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 # ================================================================================== @@ -551,6 +544,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..d5ca0a67f 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[[NDArray[np.float64]], Any] | 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 # ================================================================================== @@ -265,6 +267,11 @@ def with_step_id(self, step_id: int) -> Self: new._step_id = step_id return new + def _maybe_call_callback(self, x: NDArray[np.float64]) -> None: + """Call the optional SciPy-style ``callback(xk)`` if one was provided.""" + if self._callback is not None: + self._callback(x) + # ================================================================================== # Public attributes # ================================================================================== @@ -504,6 +511,7 @@ def _pure_evaluate_fun( exceptions=traceback, ) + self._maybe_call_callback(x) return algo_fun_value, hist_entry, log_entry def _pure_evaluate_jac( @@ -652,6 +660,7 @@ def func(x: NDArray[np.float64]) -> SpecificFunctionValue: exceptions=traceback, ) + self._maybe_call_callback(x) return (algo_fun_value, jac_value), hist_entry, log_entry def _pure_exploration_fun( @@ -786,6 +795,7 @@ def _pure_evaluate_fun_and_jac( exceptions=traceback, ) + self._maybe_call_callback(x) return (algo_fun_value, out_jac), hist_entry, log_entry diff --git a/src/optimagic/optimization/optimize.py b/src/optimagic/optimization/optimize.py index 2236269f9..5336b71c3 100644 --- a/src/optimagic/optimization/optimize.py +++ b/src/optimagic/optimization/optimize.py @@ -216,7 +216,10 @@ def maximize( args: Alternative to fun_kwargs for scipy compatibility. hess: Not yet supported. hessp: Not yet supported. - callback: Not yet supported. + callback: Optional SciPy-style callback with signature ``callback(xk)``, called + with the current internal parameter vector whenever the objective function + is evaluated. The ``callback(intermediate_result)`` interface is not yet + supported. options: Not yet supported. tol: Not yet supported. criterion: Deprecated. Use fun instead. @@ -413,7 +416,10 @@ def minimize( args: Alternative to fun_kwargs for scipy compatibility. hess: Not yet supported. hessp: Not yet supported. - callback: Not yet supported. + callback: Optional SciPy-style callback with signature ``callback(xk)``, called + with the current internal parameter vector whenever the objective function + is evaluated. The ``callback(intermediate_result)`` interface is not yet + supported. options: Not yet supported. tol: Not yet supported. criterion: Deprecated. Use fun instead. @@ -653,6 +659,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_scipy_aliases.py b/tests/optimagic/optimization/test_scipy_aliases.py index 113221674..45a54ac20 100644 --- a/tests/optimagic/optimization/test_scipy_aliases.py +++ b/tests/optimagic/optimization/test_scipy_aliases.py @@ -120,15 +120,23 @@ 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_callback_xk_is_called(): + """SciPy-style callback(xk) is invoked on objective evaluations.""" + xs = [] + + def callback(xk): + xs.append(np.asarray(xk).copy()) + + res = om.minimize( + fun=lambda x: x @ x, + x0=np.arange(3, dtype=float), + algorithm="scipy_neldermead", + callback=callback, + ) + + assert len(xs) >= 1 + assert xs[0].shape == (3,) + aaae(res.x, np.zeros(3), decimal=5) def test_exception_for_options(): From f2e5ae1f39417f87c8424572b3da0a6f3a0c2cca Mon Sep 17 00:00:00 2001 From: Pragati Date: Sun, 23 Aug 2026 18:55:00 +0530 Subject: [PATCH 2/5] ENH: make callback parallel-safe with external params (Phase 1) --- .../create_optimization_problem.py | 17 +++- .../internal_optimization_problem.py | 26 ++++-- src/optimagic/optimization/optimize.py | 26 +++--- .../optimization/test_scipy_aliases.py | 87 ++++++++++++++++++- 4 files changed, 137 insertions(+), 19 deletions(-) diff --git a/src/optimagic/optimization/create_optimization_problem.py b/src/optimagic/optimization/create_optimization_problem.py index 51eca0baa..7fd95c100 100644 --- a/src/optimagic/optimization/create_optimization_problem.py +++ b/src/optimagic/optimization/create_optimization_problem.py @@ -86,7 +86,7 @@ class OptimizationProblem: skip_checks: bool direction: Direction fun_eval: SpecificFunctionValue - callback: Callable[[Any], Any] | None + callback: Callable[[PyTree], None] | None def create_optimization_problem( @@ -522,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 # ================================================================================== diff --git a/src/optimagic/optimization/internal_optimization_problem.py b/src/optimagic/optimization/internal_optimization_problem.py index d5ca0a67f..3e7eba576 100644 --- a/src/optimagic/optimization/internal_optimization_problem.py +++ b/src/optimagic/optimization/internal_optimization_problem.py @@ -61,7 +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[[NDArray[np.float64]], Any] | None = None, + callback: Callable[[PyTree], None] | None = None, # TODO: add hess and hessp ): self._fun = fun @@ -99,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]: @@ -127,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( @@ -160,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 @@ -229,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 @@ -267,10 +273,19 @@ def with_step_id(self, step_id: int) -> Self: new._step_id = step_id return new - def _maybe_call_callback(self, x: NDArray[np.float64]) -> None: - """Call the optional SciPy-style ``callback(xk)`` if one was provided.""" + def _maybe_call_callback(self, params: PyTree) -> None: + """Call the optional SciPy-style ``callback(xk)`` if one was provided. + + Called next to history append (not inside ``_pure_*``) so it runs in the + parent process when ``n_cores > 1``. ``params`` are external user-facing + parameters (a PyTree), not the internal flat parameter vector. + + Raising ``StopIteration`` from the callback to abort optimization (as in + SciPy) is not handled yet. + + """ if self._callback is not None: - self._callback(x) + self._callback(params) # ================================================================================== # Public attributes @@ -511,7 +526,6 @@ def _pure_evaluate_fun( exceptions=traceback, ) - self._maybe_call_callback(x) return algo_fun_value, hist_entry, log_entry def _pure_evaluate_jac( @@ -660,7 +674,6 @@ def func(x: NDArray[np.float64]) -> SpecificFunctionValue: exceptions=traceback, ) - self._maybe_call_callback(x) return (algo_fun_value, jac_value), hist_entry, log_entry def _pure_exploration_fun( @@ -795,7 +808,6 @@ def _pure_evaluate_fun_and_jac( exceptions=traceback, ) - self._maybe_call_callback(x) return (algo_fun_value, out_jac), hist_entry, log_entry diff --git a/src/optimagic/optimization/optimize.py b/src/optimagic/optimization/optimize.py index 5336b71c3..946b20d89 100644 --- a/src/optimagic/optimization/optimize.py +++ b/src/optimagic/optimization/optimize.py @@ -76,8 +76,8 @@ JacType = Callable[..., PyTree] FunAndJacType = Callable[..., tuple[float | PyTree | FunctionValue, PyTree]] HessType = Callable[..., PyTree] -# TODO: refine this type -CallbackType = Callable[..., Any] +# SciPy-style callback(xk); xk is the external parameter PyTree. Returns None. +CallbackType = Callable[[PyTree], None] CriterionType = Callable[..., float | dict[str, Any]] CriterionAndDerivativeType = Callable[..., tuple[float | dict[str, Any], PyTree]] @@ -216,10 +216,13 @@ def maximize( args: Alternative to fun_kwargs for scipy compatibility. hess: Not yet supported. hessp: Not yet supported. - callback: Optional SciPy-style callback with signature ``callback(xk)``, called - with the current internal parameter vector whenever the objective function - is evaluated. The ``callback(intermediate_result)`` interface is not yet - supported. + callback: Optional SciPy-style callback with signature ``callback(xk)``, where + ``xk`` is the current **external** parameter value (a PyTree; for array + ``params`` this is a numpy array). Called next to history collection after + objective evaluations (not inside parallel worker ``_pure_*`` functions, and + not on derivative-only evaluations). Raising ``StopIteration`` to abort + optimization (as in SciPy) is not handled yet. The + ``callback(intermediate_result)`` interface is not yet supported. options: Not yet supported. tol: Not yet supported. criterion: Deprecated. Use fun instead. @@ -416,10 +419,13 @@ def minimize( args: Alternative to fun_kwargs for scipy compatibility. hess: Not yet supported. hessp: Not yet supported. - callback: Optional SciPy-style callback with signature ``callback(xk)``, called - with the current internal parameter vector whenever the objective function - is evaluated. The ``callback(intermediate_result)`` interface is not yet - supported. + callback: Optional SciPy-style callback with signature ``callback(xk)``, where + ``xk`` is the current **external** parameter value (a PyTree; for array + ``params`` this is a numpy array). Called next to history collection after + objective evaluations (not inside parallel worker ``_pure_*`` functions, and + not on derivative-only evaluations). Raising ``StopIteration`` to abort + optimization (as in SciPy) is not handled yet. The + ``callback(intermediate_result)`` interface is not yet supported. options: Not yet supported. tol: Not yet supported. criterion: Deprecated. Use fun instead. diff --git a/tests/optimagic/optimization/test_scipy_aliases.py b/tests/optimagic/optimization/test_scipy_aliases.py index 45a54ac20..a9af1de7a 100644 --- a/tests/optimagic/optimization/test_scipy_aliases.py +++ b/tests/optimagic/optimization/test_scipy_aliases.py @@ -3,7 +3,7 @@ from numpy.testing import assert_array_almost_equal as aaae import optimagic as om -from optimagic.exceptions import AliasError +from optimagic.exceptions import AliasError, InvalidFunctionError, InvalidKwargsError def test_x0_works_in_minimize(): @@ -139,6 +139,91 @@ def callback(xk): aaae(res.x, np.zeros(3), decimal=5) +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"] + + 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,) + + +def test_callback_not_called_on_jac(): + """Callback runs next to history on fun, not on jac-only evaluations.""" + from optimagic.optimization.internal_optimization_problem import ( + SphereExampleInternalOptimizationProblem, + ) + + problem = SphereExampleInternalOptimizationProblem() + calls = [] + problem._callback = lambda p: calls.append(np.asarray(p).copy()) + + x = np.ones(10) + problem.jac(x) + assert calls == [] + + problem.fun(x) + assert len(calls) == 1 + aaae(calls[0], x) + + +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", + ) + + def test_exception_for_options(): msg = "The options argument is not supported" with pytest.raises(NotImplementedError, match=msg): From 842d132382a16a345859342a916a099926e5fe33 Mon Sep 17 00:00:00 2001 From: Pragati Date: Mon, 31 Aug 2026 09:36:37 +0530 Subject: [PATCH 3/5] CI: re-run workflow after flaky macos py312 failure Co-authored-by: Cursor From 5c0b7f286ef0b92275934557260a2e3f2fca3adf Mon Sep 17 00:00:00 2001 From: Pragati Date: Sun, 6 Sep 2026 00:25:48 +0530 Subject: [PATCH 4/5] ENH: address callback review (docs + test_callback.py) --- .../internal_optimization_problem.py | 10 +- src/optimagic/optimization/optimize.py | 21 +- tests/optimagic/optimization/test_callback.py | 207 ++++++++++++++++++ .../optimization/test_scipy_aliases.py | 106 +-------- 4 files changed, 221 insertions(+), 123 deletions(-) create mode 100644 tests/optimagic/optimization/test_callback.py diff --git a/src/optimagic/optimization/internal_optimization_problem.py b/src/optimagic/optimization/internal_optimization_problem.py index 3e7eba576..0a6cb0b3b 100644 --- a/src/optimagic/optimization/internal_optimization_problem.py +++ b/src/optimagic/optimization/internal_optimization_problem.py @@ -276,12 +276,12 @@ def with_step_id(self, step_id: int) -> Self: def _maybe_call_callback(self, params: PyTree) -> None: """Call the optional SciPy-style ``callback(xk)`` if one was provided. - Called next to history append (not inside ``_pure_*``) so it runs in the - parent process when ``n_cores > 1``. ``params`` are external user-facing - parameters (a PyTree), not the internal flat parameter vector. + Args: + params: Current external (user-facing) parameters as a PyTree. - Raising ``StopIteration`` from the callback to abort optimization (as in - SciPy) is not handled yet. + Notes: + Raising ``StopIteration`` from the callback to abort optimization (as in + SciPy) is not handled yet. """ if self._callback is not None: diff --git a/src/optimagic/optimization/optimize.py b/src/optimagic/optimization/optimize.py index 946b20d89..f81c1260d 100644 --- a/src/optimagic/optimization/optimize.py +++ b/src/optimagic/optimization/optimize.py @@ -76,7 +76,6 @@ JacType = Callable[..., PyTree] FunAndJacType = Callable[..., tuple[float | PyTree | FunctionValue, PyTree]] HessType = Callable[..., PyTree] -# SciPy-style callback(xk); xk is the external parameter PyTree. Returns None. CallbackType = Callable[[PyTree], None] CriterionType = Callable[..., float | dict[str, Any]] @@ -216,12 +215,10 @@ def maximize( args: Alternative to fun_kwargs for scipy compatibility. hess: Not yet supported. hessp: Not yet supported. - callback: Optional SciPy-style callback with signature ``callback(xk)``, where - ``xk`` is the current **external** parameter value (a PyTree; for array - ``params`` this is a numpy array). Called next to history collection after - objective evaluations (not inside parallel worker ``_pure_*`` functions, and - not on derivative-only evaluations). Raising ``StopIteration`` to abort - optimization (as in SciPy) is not handled yet. The + callback: Optional callable called after each objective evaluation with + signature ``callback(xk)``, where ``xk`` is the current parameter value + (a PyTree; a numpy array if ``params`` is an array). 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. @@ -419,12 +416,10 @@ def minimize( args: Alternative to fun_kwargs for scipy compatibility. hess: Not yet supported. hessp: Not yet supported. - callback: Optional SciPy-style callback with signature ``callback(xk)``, where - ``xk`` is the current **external** parameter value (a PyTree; for array - ``params`` this is a numpy array). Called next to history collection after - objective evaluations (not inside parallel worker ``_pure_*`` functions, and - not on derivative-only evaluations). Raising ``StopIteration`` to abort - optimization (as in SciPy) is not handled yet. The + callback: Optional callable called after each objective evaluation with + signature ``callback(xk)``, where ``xk`` is the current parameter value + (a PyTree; a numpy array if ``params`` is an array). 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. diff --git a/tests/optimagic/optimization/test_callback.py b/tests/optimagic/optimization/test_callback.py new file mode 100644 index 000000000..8f3df24ab --- /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_almost_equal as aaae +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_xk_is_called(): + """SciPy-style callback(xk) is invoked on objective evaluations.""" + xs = [] + + def callback(xk): + xs.append(np.asarray(xk)) + + res = om.minimize( + fun=lambda x: x @ x, + x0=np.arange(3, dtype=float), + algorithm="scipy_neldermead", + callback=callback, + ) + + assert len(xs) >= 1 + assert xs[0].shape == (3,) + aaae(res.x, np.zeros(3), decimal=5) + + +def test_callback_matches_history_params(): + """Callback parameters match the collected optimization history.""" + xs = [] + + def callback(xk): + xs.append(np.asarray(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) + for got, expected in zip(xs, res.history.params, strict=True): + aaae(got, expected) + + +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) + aaae(got["a"], expected["a"]) + aaae(got["b"], expected["b"]) + + +def test_callback_not_called_on_jac(): + """Callback runs next to history on fun, not on jac-only evaluations.""" + from optimagic.optimization.internal_optimization_problem import ( + SphereExampleInternalOptimizationProblem, + ) + + problem = SphereExampleInternalOptimizationProblem() + calls = [] + problem._callback = lambda p: calls.append(np.asarray(p)) + + x = np.ones(10) + problem.jac(x) + assert calls == [] + + problem.fun(x) + assert len(calls) == 1 + aaae(calls[0], x) + + +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(np.asarray(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_scipy_aliases.py b/tests/optimagic/optimization/test_scipy_aliases.py index a9af1de7a..6720c5ebe 100644 --- a/tests/optimagic/optimization/test_scipy_aliases.py +++ b/tests/optimagic/optimization/test_scipy_aliases.py @@ -3,7 +3,7 @@ from numpy.testing import assert_array_almost_equal as aaae import optimagic as om -from optimagic.exceptions import AliasError, InvalidFunctionError, InvalidKwargsError +from optimagic.exceptions import AliasError def test_x0_works_in_minimize(): @@ -120,110 +120,6 @@ def test_exception_for_hessp(): ) -def test_callback_xk_is_called(): - """SciPy-style callback(xk) is invoked on objective evaluations.""" - xs = [] - - def callback(xk): - xs.append(np.asarray(xk).copy()) - - res = om.minimize( - fun=lambda x: x @ x, - x0=np.arange(3, dtype=float), - algorithm="scipy_neldermead", - callback=callback, - ) - - assert len(xs) >= 1 - assert xs[0].shape == (3,) - aaae(res.x, np.zeros(3), decimal=5) - - -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"] - - 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,) - - -def test_callback_not_called_on_jac(): - """Callback runs next to history on fun, not on jac-only evaluations.""" - from optimagic.optimization.internal_optimization_problem import ( - SphereExampleInternalOptimizationProblem, - ) - - problem = SphereExampleInternalOptimizationProblem() - calls = [] - problem._callback = lambda p: calls.append(np.asarray(p).copy()) - - x = np.ones(10) - problem.jac(x) - assert calls == [] - - problem.fun(x) - assert len(calls) == 1 - aaae(calls[0], x) - - -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", - ) - - def test_exception_for_options(): msg = "The options argument is not supported" with pytest.raises(NotImplementedError, match=msg): From 278131667bf7a07aaec14f8e5eafeaf170703191 Mon Sep 17 00:00:00 2001 From: Janos Gabler Date: Fri, 25 Sep 2026 14:19:41 +0200 Subject: [PATCH 5/5] Address review comments on callback support - Make callback a required argument of InternalOptimizationProblem. - Document that xk is not copied and that the callback is not called during multistart exploration. - Update the SciPy alignment enhancement proposal. - Compare callback params against history exactly, drop the redundant test, and test the jac case through the public API. Co-Authored-By: Claude Opus 5.5 --- docs/source/development/ep-03-alignment.md | 8 +- .../internal_optimization_problem.py | 4 +- src/optimagic/optimization/optimize.py | 16 ++-- tests/optimagic/optimization/test_callback.py | 80 +++++++++---------- .../test_internal_optimization_problem.py | 3 + 5 files changed, 61 insertions(+), 50 deletions(-) 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/internal_optimization_problem.py b/src/optimagic/optimization/internal_optimization_problem.py index 0a6cb0b3b..9c7b52b04 100644 --- a/src/optimagic/optimization/internal_optimization_problem.py +++ b/src/optimagic/optimization/internal_optimization_problem.py @@ -61,7 +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 = None, + callback: Callable[[PyTree], None] | None, # TODO: add hess and hessp ): self._fun = fun @@ -980,6 +980,7 @@ def __init__( linear_constraints=linear_constraints, nonlinear_constraints=nonlinear_constraints, logger=logger, + callback=None, ) @@ -1118,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 f81c1260d..a2c2ecd5d 100644 --- a/src/optimagic/optimization/optimize.py +++ b/src/optimagic/optimization/optimize.py @@ -216,9 +216,11 @@ def maximize( hess: Not yet supported. hessp: Not yet supported. callback: Optional callable called after each objective evaluation with - signature ``callback(xk)``, where ``xk`` is the current parameter value - (a PyTree; a numpy array if ``params`` is an array). Raising - ``StopIteration`` to abort optimization is not yet supported. The + 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. @@ -417,9 +419,11 @@ def minimize( hess: Not yet supported. hessp: Not yet supported. callback: Optional callable called after each objective evaluation with - signature ``callback(xk)``, where ``xk`` is the current parameter value - (a PyTree; a numpy array if ``params`` is an array). Raising - ``StopIteration`` to abort optimization is not yet supported. The + 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. diff --git a/tests/optimagic/optimization/test_callback.py b/tests/optimagic/optimization/test_callback.py index 8f3df24ab..e2d2a3d9d 100644 --- a/tests/optimagic/optimization/test_callback.py +++ b/tests/optimagic/optimization/test_callback.py @@ -4,7 +4,6 @@ import numpy as np import pytest -from numpy.testing import assert_array_almost_equal as aaae from numpy.testing import assert_array_equal as aae import optimagic as om @@ -15,31 +14,12 @@ from optimagic.typing import AggregationLevel -def test_callback_xk_is_called(): - """SciPy-style callback(xk) is invoked on objective evaluations.""" - xs = [] - - def callback(xk): - xs.append(np.asarray(xk)) - - res = om.minimize( - fun=lambda x: x @ x, - x0=np.arange(3, dtype=float), - algorithm="scipy_neldermead", - callback=callback, - ) - - assert len(xs) >= 1 - assert xs[0].shape == (3,) - aaae(res.x, np.zeros(3), decimal=5) - - def test_callback_matches_history_params(): """Callback parameters match the collected optimization history.""" xs = [] def callback(xk): - xs.append(np.asarray(xk)) + xs.append(xk) res = om.minimize( fun=lambda x: x @ x, @@ -50,8 +30,7 @@ def callback(xk): assert res.history is not None assert len(xs) == len(res.history.params) - for got, expected in zip(xs, res.history.params, strict=True): - aaae(got, expected) + aae(xs, res.history.params) def test_callback_receives_external_params(): @@ -83,27 +62,48 @@ def fun(p): assert len(received) == len(res.history.params) for got, expected in zip(received, res.history.params, strict=True): assert set(got) == set(expected) - aaae(got["a"], expected["a"]) - aaae(got["b"], expected["b"]) + aae(got["a"], expected["a"]) + aae(got["b"], expected["b"]) -def test_callback_not_called_on_jac(): - """Callback runs next to history on fun, not on jac-only evaluations.""" - from optimagic.optimization.internal_optimization_problem import ( - SphereExampleInternalOptimizationProblem, - ) +@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) - problem = SphereExampleInternalOptimizationProblem() - calls = [] - problem._callback = lambda p: calls.append(np.asarray(p)) - x = np.ones(10) - problem.jac(x) - assert calls == [] +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, + ) - problem.fun(x) - assert len(calls) == 1 - aaae(calls[0], x) + aae(xs, [np.arange(3) + 2.0]) + assert len(res.history.params) == 3 def test_invalid_callback_too_few_arguments(): @@ -192,7 +192,7 @@ def test_callback_history_with_parallel_optimizer(): collected = [] def callback(xk): - collected.append(np.asarray(xk)) + collected.append(xk) res = minimize( fun=lambda x: 5.0, 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