Skip to content

Commit 7a4d5cc

Browse files
committed
feat: adapt solutions to new initial guess grids
1 parent 10fb30c commit 7a4d5cc

4 files changed

Lines changed: 172 additions & 0 deletions

File tree

bioptim/__init__.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -249,6 +249,7 @@
249249
)
250250
from .optimization.receding_horizon_optimization import MovingHorizonEstimator, NonlinearModelPredictiveControl
251251
from .optimization.solution.solution import Solution
252+
from .optimization.solution.utils import adapt_solution_to_initial_guesses
252253
from .optimization.solution.solution_data import SolutionMerge, TimeAlignment
253254
from .optimization.stochastic_optimal_control_program import StochasticOptimalControlProgram
254255
from .optimization.variable_scaling import VariableScalingList, VariableScaling

bioptim/optimization/solution/solution.py

Lines changed: 19 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -736,6 +736,25 @@ def copy(self, skip_data: Bool = False) -> "Solution":
736736
new._parameters = deepcopy(self._parameters)
737737
return new
738738

739+
def to_initial_guesses(
740+
self,
741+
state_initial_guesses: InitialGuessList,
742+
control_initial_guesses: InitialGuessList,
743+
parameter_initial_guesses: InitialGuessList | None = None,
744+
algebraic_state_initial_guesses: InitialGuessList | None = None,
745+
) -> tuple[InitialGuessList, InitialGuessList, InitialGuessList, InitialGuessList]:
746+
"""Adapt this solution's primal variables to new initial-guess grids."""
747+
748+
from .utils import adapt_solution_to_initial_guesses
749+
750+
return adapt_solution_to_initial_guesses(
751+
self,
752+
state_initial_guesses,
753+
control_initial_guesses,
754+
parameter_initial_guesses,
755+
algebraic_state_initial_guesses,
756+
)
757+
739758
def _prepare_integrate(self, integrator: SolutionIntegrator) -> AnyTuple:
740759
"""
741760
Prepare the variables for the states integration and checks if the integrator is compatible with the ocp.

bioptim/optimization/solution/utils.py

Lines changed: 99 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,105 @@
66
NpArrayDict,
77
)
88

9+
from ...limits.path_conditions import InitialGuessList
10+
from ...misc.enums import InterpolationType
11+
from .solution_data import SolutionMerge
12+
13+
14+
def _resample_initial_guess(values: np.ndarray, n_columns: int) -> np.ndarray:
15+
"""Linearly resample node values while preserving both endpoints."""
16+
17+
values = np.asarray(values, dtype=float)
18+
if values.ndim == 1:
19+
values = values[:, np.newaxis]
20+
finite_columns = np.all(np.isfinite(values), axis=0)
21+
values = values[:, finite_columns]
22+
if values.shape[1] == 0:
23+
raise ValueError("A solution variable contains no finite node that can be transferred")
24+
if n_columns == 1:
25+
return values[:, :1].copy()
26+
if values.shape[1] == 1:
27+
return np.repeat(values, n_columns, axis=1)
28+
29+
source_grid = np.linspace(0.0, 1.0, values.shape[1])
30+
target_grid = np.linspace(0.0, 1.0, n_columns)
31+
return np.vstack([np.interp(target_grid, source_grid, row) for row in values])
32+
33+
34+
def _target_column_count(initial_guess) -> int:
35+
interpolation = initial_guess.type
36+
if interpolation == InterpolationType.CONSTANT:
37+
return 1
38+
if interpolation == InterpolationType.CONSTANT_WITH_FIRST_AND_LAST_DIFFERENT:
39+
return 3
40+
if interpolation == InterpolationType.LINEAR:
41+
return 2
42+
if interpolation in (InterpolationType.EACH_FRAME, InterpolationType.ALL_POINTS):
43+
return initial_guess.init.shape[1]
44+
raise NotImplementedError(f"Adapting a solution to {interpolation} is not implemented")
45+
46+
47+
def _as_phase_list(data, n_phases: int) -> list[dict]:
48+
if n_phases == 1 and isinstance(data, dict):
49+
return [data]
50+
if not isinstance(data, list) or len(data) != n_phases:
51+
raise ValueError(f"The solution has {len(data) if isinstance(data, list) else 1} phases, expected {n_phases}")
52+
return data
53+
54+
55+
def _adapt_variable_group(source, target: InitialGuessList, group_name: str) -> InitialGuessList:
56+
adapted = InitialGuessList()
57+
n_phases = len(target.options)
58+
source_phases = _as_phase_list(source, n_phases)
59+
for phase, target_phase in enumerate(target.options):
60+
for key, target_guess in target_phase.items():
61+
if key not in source_phases[phase]:
62+
raise KeyError(f"{group_name} '{key}' is absent from the solution phase {phase}")
63+
values = _resample_initial_guess(source_phases[phase][key], _target_column_count(target_guess))
64+
adapted.add(key, values, interpolation=target_guess.type, phase=phase)
65+
return adapted
66+
67+
68+
def adapt_solution_to_initial_guesses(
69+
solution,
70+
state_initial_guesses: InitialGuessList,
71+
control_initial_guesses: InitialGuessList,
72+
parameter_initial_guesses: InitialGuessList | None = None,
73+
algebraic_state_initial_guesses: InitialGuessList | None = None,
74+
) -> tuple[InitialGuessList, InitialGuessList, InitialGuessList, InitialGuessList]:
75+
"""Adapt a solution's primal variables to the grids described by new initial-guess lists.
76+
77+
Resampling is performed on a normalized phase grid, so it supports different numbers of shooting nodes,
78+
collocation points and control nodes. Solver multipliers are deliberately not transferred by this function.
79+
"""
80+
81+
states = _adapt_variable_group(
82+
solution.decision_states(to_merge=SolutionMerge.NODES), state_initial_guesses, "State"
83+
)
84+
controls = _adapt_variable_group(
85+
solution.decision_controls(to_merge=SolutionMerge.NODES), control_initial_guesses, "Control"
86+
)
87+
88+
parameters = InitialGuessList()
89+
if parameter_initial_guesses is not None:
90+
source_parameters = solution.decision_parameters()
91+
for phase, target_phase in enumerate(parameter_initial_guesses.options):
92+
for key, target_guess in target_phase.items():
93+
if key not in source_parameters:
94+
raise KeyError(f"Parameter '{key}' is absent from the solution")
95+
values = np.asarray(source_parameters[key], dtype=float).reshape((-1, 1))
96+
parameters.add(key, values, interpolation=target_guess.type, phase=phase)
97+
98+
algebraic_states = InitialGuessList()
99+
if algebraic_state_initial_guesses is not None:
100+
algebraic_states = _adapt_variable_group(
101+
solution.decision_algebraic_states(to_merge=SolutionMerge.NODES),
102+
algebraic_state_initial_guesses,
103+
"Algebraic state",
104+
)
105+
106+
return states, controls, parameters, algebraic_states
107+
9108

10109
def concatenate_optimization_variables_dict(variable: list[NpArrayDict], continuous: Bool = True) -> list[NpArrayDict]:
11110
"""
Lines changed: 53 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,53 @@
1+
import numpy as np
2+
3+
from bioptim import InitialGuessList, InterpolationType, adapt_solution_to_initial_guesses
4+
5+
6+
class _FakeSolution:
7+
def decision_states(self, **_):
8+
return {"q": np.array([[0.0, 1.0, 2.0]])}
9+
10+
def decision_controls(self, **_):
11+
return {"tau": np.array([[0.0, 2.0, np.nan]])}
12+
13+
def decision_parameters(self):
14+
return {"mass": np.array([3.0])}
15+
16+
def decision_algebraic_states(self, **_):
17+
return {"contact": np.array([[1.0, 3.0, 5.0]])}
18+
19+
20+
def test_adapt_solution_between_different_grids():
21+
x_init = InitialGuessList()
22+
x_init.add("q", np.zeros((1, 5)), interpolation=InterpolationType.EACH_FRAME)
23+
u_init = InitialGuessList()
24+
u_init.add("tau", np.zeros((1, 4)), interpolation=InterpolationType.EACH_FRAME)
25+
p_init = InitialGuessList()
26+
p_init.add("mass", [0.0], interpolation=InterpolationType.CONSTANT)
27+
a_init = InitialGuessList()
28+
a_init.add("contact", np.zeros((1, 7)), interpolation=InterpolationType.ALL_POINTS)
29+
30+
states, controls, parameters, algebraic_states = adapt_solution_to_initial_guesses(
31+
_FakeSolution(), x_init, u_init, p_init, a_init
32+
)
33+
34+
np.testing.assert_allclose(states[0]["q"].init, [[0.0, 0.5, 1.0, 1.5, 2.0]])
35+
np.testing.assert_allclose(controls[0]["tau"].init, [[0.0, 2 / 3, 4 / 3, 2.0]])
36+
np.testing.assert_allclose(parameters[0]["mass"].init, [[3.0]])
37+
np.testing.assert_allclose(algebraic_states[0]["contact"].init, [[1, 5 / 3, 7 / 3, 3, 11 / 3, 13 / 3, 5]])
38+
39+
40+
def test_adapt_solution_preserves_control_grid_semantics():
41+
x_init = InitialGuessList()
42+
x_init.add("q", np.zeros((1, 3)), interpolation=InterpolationType.EACH_FRAME)
43+
u_init = InitialGuessList()
44+
u_init.add("tau", np.zeros((1, 2)), interpolation=InterpolationType.LINEAR)
45+
46+
_, controls, _, _ = _FakeSolutionAdapter().to_initial_guesses(x_init, u_init)
47+
48+
assert controls[0]["tau"].type == InterpolationType.LINEAR
49+
np.testing.assert_allclose(controls[0]["tau"].init, [[0.0, 2.0]])
50+
51+
52+
class _FakeSolutionAdapter(_FakeSolution):
53+
to_initial_guesses = __import__("bioptim").Solution.to_initial_guesses

0 commit comments

Comments
 (0)