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
8 changes: 5 additions & 3 deletions docs/source/development/ep-03-alignment.md
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand Down
25 changes: 17 additions & 8 deletions src/optimagic/optimization/create_optimization_problem.py
Original file line number Diff line number Diff line change
Expand Up @@ -86,6 +86,7 @@ class OptimizationProblem:
skip_checks: bool
direction: Direction
fun_eval: SpecificFunctionValue
callback: Callable[[PyTree], None] | None


def create_optimization_problem(
Expand Down Expand Up @@ -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
# ==================================================================================
Expand Down Expand Up @@ -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
# ==================================================================================
Expand All @@ -551,6 +559,7 @@ def create_optimization_problem(
skip_checks=skip_checks,
direction=direction,
fun_eval=fun_eval,
callback=callback,

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Currently you don't do any validation of the user passed callback. For example, if the user passes a callable with the wrong signature, this would lead to a hard to understand error in the first iteration. We need the same validation for callbacks as we have for all other user evaluated functions like fun, jac, etc.

We also need tests that invalid callbacks are handled well.

)

return problem
Expand Down
24 changes: 24 additions & 0 deletions src/optimagic/optimization/internal_optimization_problem.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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

# ==================================================================================
Expand All @@ -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]:
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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

Expand Down Expand Up @@ -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

Expand Down Expand Up @@ -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
# ==================================================================================
Expand Down Expand Up @@ -958,6 +980,7 @@ def __init__(
linear_constraints=linear_constraints,
nonlinear_constraints=nonlinear_constraints,
logger=logger,
callback=None,
)


Expand Down Expand Up @@ -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,
)
20 changes: 16 additions & 4 deletions src/optimagic/optimization/optimize.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]]
Expand Down Expand Up @@ -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.
Expand Down Expand Up @@ -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.
Expand Down Expand Up @@ -653,6 +664,7 @@ def _optimize(problem: OptimizationProblem) -> OptimizeResult:
linear_constraints=None,
nonlinear_constraints=internal_nonlinear_constraints,
logger=logger,
callback=problem.callback,
)

# ==================================================================================
Expand Down
Loading
Loading