Skip to content
Draft
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
81 changes: 81 additions & 0 deletions bioptim/optimization/receding_horizon_optimization.py
Original file line number Diff line number Diff line change
Expand Up @@ -296,13 +296,15 @@ def advance_window(self, sol: Solution, steps: Int = 0, **advance_options) -> No

init_states_have_changed = self.advance_window_initial_guess_states(sol, **advance_options)
init_controls_have_changed = self.advance_window_initial_guess_controls(sol, **advance_options)
init_algebraic_states_have_changed = self.advance_window_initial_guess_algebraic_states(sol, **advance_options)
init_parameter_have_changed = self.advance_window_initial_guess_parameters(sol, **advance_options)

if self.ocp_solver.opts.type != SolverType.ACADOS:
self.update_initial_guess(
self.nlp[0].x_init if init_states_have_changed else None,
self.nlp[0].u_init if init_controls_have_changed else None,
self.parameter_init if init_parameter_have_changed else None,
self.nlp[0].a_init if init_algebraic_states_have_changed else None,
)

def advance_window_bounds_states(self, sol: Solution, **advance_options) -> Bool:
Expand Down Expand Up @@ -366,6 +368,26 @@ def advance_window_initial_guess_controls(self, sol: Solution, **advance_options
)
return True

def advance_window_initial_guess_algebraic_states(self, sol: Solution, **advance_options) -> Bool:
algebraic_states = sol.decision_algebraic_states(to_merge=SolutionMerge.NODES)

for key in algebraic_states.keys():
if self.nlp[0].a_init[key].type != InterpolationType.EACH_FRAME:
self.nlp[0].a_init.add(
key,
np.ndarray(algebraic_states[key].shape),
interpolation=InterpolationType.EACH_FRAME,
phase=0,
)
self.nlp[0].a_init[key].check_and_adjust_dimensions(
len(self.nlp[0].algebraic_states[key]), self.nlp[0].n_algebraic_states_nodes - 1
)

self.nlp[0].a_init[key].init[:, :] = np.concatenate(
(algebraic_states[key][:, 1:], algebraic_states[key][:, -1][:, np.newaxis]), axis=1
)
return bool(algebraic_states)

def advance_window_initial_guess_parameters(self, sol: Solution, **advance_options) -> Bool:
parameters = sol.parameters
for key in parameters.keys():
Expand Down Expand Up @@ -636,6 +658,24 @@ def advance_window_initial_guess_controls(self, sol: Solution, **advance_options
self.nlp[0].u_init[key].init[:, :] = controls[key][:, :]
return True

def advance_window_initial_guess_algebraic_states(self, sol: Solution, **advance_options) -> Bool:
algebraic_states = sol.decision_algebraic_states(to_merge=SolutionMerge.NODES)

for key in algebraic_states.keys():
if self.nlp[0].a_init[key].type != InterpolationType.EACH_FRAME:
self.nlp[0].a_init.add(
key,
np.ndarray(algebraic_states[key].shape),
interpolation=InterpolationType.EACH_FRAME,
phase=0,
)
self.nlp[0].a_init[key].check_and_adjust_dimensions(
len(self.nlp[0].algebraic_states[key]), self.nlp[0].n_algebraic_states_nodes - 1
)

self.nlp[0].a_init[key].init[:, :] = algebraic_states[key]
return bool(algebraic_states)


class MultiCyclicRecedingHorizonOptimization(CyclicRecedingHorizonOptimization):
def __init__(
Expand Down Expand Up @@ -751,6 +791,47 @@ def advance_window_initial_guess_controls(self, sol: Solution, **advance_options
raise NotImplementedError(f"Control type {self.nlp[0].control_type} is not implemented yet")
self.nlp[0].u_init[key].init[:, :] = controls[key][:, frames]

def advance_window_initial_guess_algebraic_states(self, sol: Solution, **advance_options) -> Bool:
algebraic_states = sol.decision_algebraic_states(to_merge=SolutionMerge.NODES)

for key in algebraic_states.keys():
if isinstance(self.nlp[0].dynamics_type.ode_solver, OdeSolver.COLLOCATION):
if self.nlp[0].a_init[key].type != InterpolationType.ALL_POINTS:
self.nlp[0].a_init.add(
key,
np.ndarray((algebraic_states[key].shape[0], self.nlp[0].ns * self.nb_intermediate_frames + 1)),
interpolation=InterpolationType.ALL_POINTS,
phase=0,
)
self.nlp[0].a_init[key].check_and_adjust_dimensions(
self.nlp[0].algebraic_states[key].shape,
self.nlp[0].ns * self.nb_intermediate_frames,
)
frames = []
for _ in range(self.n_cycles):
frames.extend(
range(
self.n_cycles_to_advance * self.cycle_len * self.nb_intermediate_frames,
(self.n_cycles_to_advance + 1) * self.cycle_len * self.nb_intermediate_frames,
)
)
frames.append((self.n_cycles_to_advance + 1) * self.cycle_len * self.nb_intermediate_frames)
else:
if self.nlp[0].a_init[key].type != InterpolationType.EACH_FRAME:
self.nlp[0].a_init.add(
key,
np.ndarray((algebraic_states[key].shape[0], self.nlp[0].ns + 1)),
interpolation=InterpolationType.EACH_FRAME,
phase=0,
)
self.nlp[0].a_init[key].check_and_adjust_dimensions(
self.nlp[0].algebraic_states[key].shape, self.nlp[0].ns
)
frames = self.initial_guess_frames

self.nlp[0].a_init[key].init[:, :] = algebraic_states[key][:, frames]
return bool(algebraic_states)

def solve(
self,
update_function: Callable | None = None,
Expand Down
21 changes: 21 additions & 0 deletions tests/shard1/test_mhe.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@
import os
import shutil
from sys import platform
from types import SimpleNamespace

from bioptim.misc.enums import SolverType
from bioptim import (
Expand All @@ -20,6 +21,26 @@
from ..utils import TestUtils


def test_mhe_advances_algebraic_state_initial_guesses():
mhe = MovingHorizonEstimator.__new__(MovingHorizonEstimator)
algebraic_init = SimpleNamespace(
type=InterpolationType.EACH_FRAME,
init=np.zeros((1, 4)),
)
mhe.nlp = [SimpleNamespace(a_init={"lambda": algebraic_init})]

class Solution:
@staticmethod
def decision_algebraic_states(to_merge):
assert to_merge == SolutionMerge.NODES
return {"lambda": np.array([[1.0, 2.0, 3.0, 4.0]])}

changed = mhe.advance_window_initial_guess_algebraic_states(Solution())

assert changed
npt.assert_equal(algebraic_init.init, [[2.0, 3.0, 4.0, 4.0]])


@pytest.mark.parametrize("phase_dynamics", [PhaseDynamics.SHARED_DURING_THE_PHASE, PhaseDynamics.ONE_PER_NODE])
@pytest.mark.parametrize("solver", [Solver.IPOPT, Solver.ACADOS])
def test_mhe(solver, phase_dynamics):
Expand Down
Loading