|
6 | 6 | NpArrayDict, |
7 | 7 | ) |
8 | 8 |
|
| 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 | + |
9 | 108 |
|
10 | 109 | def concatenate_optimization_variables_dict(variable: list[NpArrayDict], continuous: Bool = True) -> list[NpArrayDict]: |
11 | 110 | """ |
|
0 commit comments