1010 lines
36 KiB
Python
1010 lines
36 KiB
Python
from __future__ import annotations
|
|
|
|
from app.simulation.core.errors import RecoverableTrialStateError
|
|
|
|
import math
|
|
from dataclasses import dataclass
|
|
from typing import Callable, Literal, Sequence
|
|
|
|
|
|
CancellationCheck = Callable[[], bool]
|
|
AcceptedStepCallback = Callable[[float], None]
|
|
IntegrationStatus = Literal["completed", "cancelled", "failed"]
|
|
DenseState = Callable[[float], list[float]]
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class StateTransition:
|
|
"""A state reset located inside an accepted integration step."""
|
|
|
|
time: float
|
|
state: list[float]
|
|
|
|
|
|
StateTransitionHandler = Callable[
|
|
[float, list[float], float, list[float], DenseState],
|
|
StateTransition | None,
|
|
]
|
|
|
|
_MAX_STATE_TRANSITIONS_AT_SAME_TIME = 64
|
|
|
|
|
|
class _IntegrationCancelled(Exception):
|
|
pass
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class SolveIVPConfig:
|
|
t_start: float = 0.0
|
|
t_stop: float = 20.0
|
|
method: str = "BDF"
|
|
rtol: float = 1e-6
|
|
atol: float = 1e-8
|
|
max_step: float = 1e-3
|
|
first_step: float | None = None
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class ODESolution:
|
|
t: list[float]
|
|
y: list[list[float]]
|
|
success: bool
|
|
message: str
|
|
status: IntegrationStatus = "completed"
|
|
error: Exception | None = None
|
|
|
|
|
|
def _vector_add(a: list[float], b: list[float], scale: float = 1.0) -> list[float]:
|
|
return [x + scale * y for x, y in zip(a, b)]
|
|
|
|
|
|
def _append_solution_sample(
|
|
times: list[float],
|
|
states: list[list[float]],
|
|
time: float,
|
|
state: list[float],
|
|
) -> None:
|
|
time = float(time)
|
|
if times and time <= times[-1]:
|
|
return
|
|
times.append(time)
|
|
for index, value in enumerate(state):
|
|
states[index].append(float(value))
|
|
|
|
|
|
def _append_or_replace_solution_sample(
|
|
times: list[float],
|
|
states: list[list[float]],
|
|
time: float,
|
|
state: list[float],
|
|
) -> None:
|
|
"""Store a reset state even when its event time was already sampled."""
|
|
|
|
time = float(time)
|
|
if times and time == times[-1]:
|
|
times[-1] = time
|
|
for index, value in enumerate(state):
|
|
states[index][-1] = float(value)
|
|
return
|
|
_append_solution_sample(times, states, time, state)
|
|
|
|
|
|
def _normalize_state_transition(
|
|
transition: StateTransition,
|
|
before_time: float,
|
|
after_time: float,
|
|
state_size: int,
|
|
) -> StateTransition:
|
|
"""Validate and normalize a transition returned for an accepted step."""
|
|
|
|
if not isinstance(transition, StateTransition):
|
|
raise TypeError(
|
|
"State transition handlers must return StateTransition or None."
|
|
)
|
|
|
|
transition_time = float(transition.time)
|
|
if not math.isfinite(transition_time):
|
|
raise ValueError("State transition times must be finite numbers.")
|
|
|
|
tolerance = 16.0 * max(
|
|
math.ulp(before_time),
|
|
math.ulp(after_time),
|
|
math.ulp(transition_time),
|
|
)
|
|
if (
|
|
transition_time < before_time - tolerance
|
|
or transition_time > after_time + tolerance
|
|
):
|
|
raise ValueError(
|
|
"State transition time must lie inside the accepted integration step."
|
|
)
|
|
transition_time = min(max(transition_time, before_time), after_time)
|
|
|
|
transition_state = [float(value) for value in transition.state]
|
|
if len(transition_state) != state_size:
|
|
raise ValueError(
|
|
"State transition reset state must have the same size as the ODE state."
|
|
)
|
|
if not all(math.isfinite(value) for value in transition_state):
|
|
raise ValueError("State transition reset states must contain finite numbers.")
|
|
return StateTransition(time=transition_time, state=transition_state)
|
|
|
|
|
|
def _is_repeated_state_transition(
|
|
transition: StateTransition,
|
|
last_transition: StateTransition | None,
|
|
) -> bool:
|
|
"""Suppress only the exact reset that was just applied.
|
|
|
|
A second reset at the same instant is meaningful when it produces a
|
|
different state (for example, two constraints becoming active together).
|
|
"""
|
|
|
|
return (
|
|
last_transition is not None
|
|
and transition.time == last_transition.time
|
|
and transition.state == last_transition.state
|
|
)
|
|
|
|
|
|
def _next_same_time_transition_count(
|
|
transition: StateTransition,
|
|
last_transition: StateTransition | None,
|
|
previous_count: int,
|
|
) -> int:
|
|
count = (
|
|
previous_count + 1
|
|
if last_transition is not None
|
|
and transition.time == last_transition.time
|
|
else 1
|
|
)
|
|
if count > _MAX_STATE_TRANSITIONS_AT_SAME_TIME:
|
|
raise RuntimeError(
|
|
"State transition handler exceeded "
|
|
f"{_MAX_STATE_TRANSITIONS_AT_SAME_TIME} chained resets at the same time."
|
|
)
|
|
return count
|
|
|
|
|
|
def _align_transition_with_exact_endpoint(
|
|
transition: StateTransition,
|
|
requested_time: float,
|
|
exact_endpoint: float | None,
|
|
) -> StateTransition:
|
|
"""Keep an event reported at a breakpoint on that exact public timestamp."""
|
|
|
|
if exact_endpoint is not None and requested_time == exact_endpoint:
|
|
return StateTransition(
|
|
time=float(exact_endpoint),
|
|
state=list(transition.state),
|
|
)
|
|
return transition
|
|
|
|
|
|
def _normalize_breakpoints(
|
|
config: SolveIVPConfig,
|
|
breakpoints: Sequence[float] | None,
|
|
) -> list[float]:
|
|
"""Return sorted, unique breakpoints strictly inside the integration span."""
|
|
|
|
if breakpoints is None or len(breakpoints) == 0:
|
|
return []
|
|
if config.t_stop < config.t_start:
|
|
raise ValueError("Segmented integration requires t_stop to follow t_start.")
|
|
|
|
normalized: list[float] = []
|
|
for raw_breakpoint in breakpoints:
|
|
breakpoint = float(raw_breakpoint)
|
|
if not math.isfinite(breakpoint):
|
|
raise ValueError("Integration breakpoints must be finite numbers.")
|
|
if config.t_start < breakpoint < config.t_stop:
|
|
normalized.append(breakpoint)
|
|
|
|
normalized.sort()
|
|
return [
|
|
breakpoint
|
|
for index, breakpoint in enumerate(normalized)
|
|
if index == 0 or breakpoint != normalized[index - 1]
|
|
]
|
|
|
|
|
|
def _runge_kutta_4(
|
|
rhs: Callable[[float, list[float]], list[float]],
|
|
initial_state: list[float],
|
|
config: SolveIVPConfig,
|
|
t_eval: list[float] | None,
|
|
cancel_check: CancellationCheck | None = None,
|
|
accepted_step_callback: AcceptedStepCallback | None = None,
|
|
state_transition_handler: StateTransitionHandler | None = None,
|
|
) -> ODESolution:
|
|
if t_eval is None:
|
|
point_count = max(
|
|
2,
|
|
int((config.t_stop - config.t_start) / max(config.max_step, 1e-6)) + 1,
|
|
)
|
|
step = (config.t_stop - config.t_start) / (point_count - 1)
|
|
t_eval = [config.t_start + index * step for index in range(point_count)]
|
|
|
|
state = list(initial_state)
|
|
states = [[value] for value in state]
|
|
times = [float(t_eval[0])]
|
|
current_time = float(t_eval[0])
|
|
status: IntegrationStatus = "completed"
|
|
message = "Integrated with built-in RK4 fallback because SciPy is unavailable."
|
|
error: Exception | None = None
|
|
last_transition: StateTransition | None = None
|
|
same_time_transition_count = 0
|
|
last_reported_step: float | None = None
|
|
|
|
def report_step(time: float) -> None:
|
|
nonlocal last_reported_step
|
|
if accepted_step_callback is None:
|
|
return
|
|
if last_reported_step is not None and time <= last_reported_step:
|
|
return
|
|
accepted_step_callback(float(time))
|
|
last_reported_step = float(time)
|
|
|
|
try:
|
|
for target_time in t_eval[1:]:
|
|
while current_time < target_time:
|
|
if cancel_check is not None and cancel_check():
|
|
raise _IntegrationCancelled
|
|
dt = min(config.max_step, target_time - current_time)
|
|
k1 = rhs(current_time, state)
|
|
k2 = rhs(current_time + 0.5 * dt, _vector_add(state, k1, 0.5 * dt))
|
|
k3 = rhs(current_time + 0.5 * dt, _vector_add(state, k2, 0.5 * dt))
|
|
k4 = rhs(current_time + dt, _vector_add(state, k3, dt))
|
|
next_state = [
|
|
value + (dt / 6.0) * (a + 2.0 * b + 2.0 * c + d)
|
|
for value, a, b, c, d in zip(state, k1, k2, k3, k4)
|
|
]
|
|
next_time = current_time + dt
|
|
|
|
transition: StateTransition | None = None
|
|
if state_transition_handler is not None:
|
|
step_start = current_time
|
|
step_state = list(state)
|
|
|
|
def dense_state(time: float) -> list[float]:
|
|
fraction = (float(time) - step_start) / (next_time - step_start)
|
|
return [
|
|
before + fraction * (after - before)
|
|
for before, after in zip(step_state, next_state)
|
|
]
|
|
|
|
candidate = state_transition_handler(
|
|
step_start,
|
|
list(step_state),
|
|
next_time,
|
|
list(next_state),
|
|
dense_state,
|
|
)
|
|
if candidate is not None:
|
|
candidate = _normalize_state_transition(
|
|
candidate,
|
|
step_start,
|
|
next_time,
|
|
len(state),
|
|
)
|
|
if not _is_repeated_state_transition(
|
|
candidate,
|
|
last_transition,
|
|
):
|
|
transition = candidate
|
|
|
|
if transition is not None:
|
|
same_time_transition_count = _next_same_time_transition_count(
|
|
transition,
|
|
last_transition,
|
|
same_time_transition_count,
|
|
)
|
|
current_time = transition.time
|
|
state = list(transition.state)
|
|
last_transition = transition
|
|
_append_or_replace_solution_sample(
|
|
times,
|
|
states,
|
|
current_time,
|
|
state,
|
|
)
|
|
else:
|
|
current_time = next_time
|
|
state = next_state
|
|
report_step(current_time)
|
|
|
|
_append_solution_sample(times, states, target_time, state)
|
|
except _IntegrationCancelled:
|
|
status = "cancelled"
|
|
message = "Simulation was stopped before reaching the requested end time."
|
|
_append_solution_sample(times, states, current_time, state)
|
|
except Exception as exc:
|
|
status = "failed"
|
|
message = str(exc)
|
|
error = exc
|
|
_append_solution_sample(times, states, current_time, state)
|
|
|
|
return ODESolution(
|
|
t=times,
|
|
y=states,
|
|
success=status == "completed",
|
|
message=message,
|
|
status=status,
|
|
error=error,
|
|
)
|
|
|
|
|
|
def _runge_kutta_4_segmented(
|
|
rhs: Callable[[float, list[float]], list[float]],
|
|
initial_state: list[float],
|
|
config: SolveIVPConfig,
|
|
t_eval: list[float] | None,
|
|
breakpoints: Sequence[float],
|
|
cancel_check: CancellationCheck | None = None,
|
|
accepted_step_callback: AcceptedStepCallback | None = None,
|
|
state_transition_handler: StateTransitionHandler | None = None,
|
|
) -> ODESolution:
|
|
"""RK4 fallback that never evaluates a pre-breakpoint step at the breakpoint."""
|
|
|
|
if t_eval is None:
|
|
point_count = max(
|
|
2,
|
|
int((config.t_stop - config.t_start) / max(config.max_step, 1e-6)) + 1,
|
|
)
|
|
sample_step = (config.t_stop - config.t_start) / (point_count - 1)
|
|
sample_times = [
|
|
config.t_start + index * sample_step for index in range(point_count)
|
|
]
|
|
else:
|
|
sample_times = [float(time) for time in t_eval]
|
|
|
|
state = [float(value) for value in initial_state]
|
|
states = [[value] for value in state]
|
|
times = [float(config.t_start)]
|
|
current_time = float(config.t_start)
|
|
sample_index = 0
|
|
while (
|
|
sample_index < len(sample_times)
|
|
and sample_times[sample_index] <= config.t_start
|
|
):
|
|
sample_index += 1
|
|
|
|
status: IntegrationStatus = "completed"
|
|
message = "Integrated with built-in RK4 fallback because SciPy is unavailable."
|
|
error: Exception | None = None
|
|
last_transition: StateTransition | None = None
|
|
same_time_transition_count = 0
|
|
last_reported_step: float | None = None
|
|
|
|
def report_step(time: float) -> None:
|
|
nonlocal last_reported_step
|
|
if accepted_step_callback is None:
|
|
return
|
|
if last_reported_step is not None and time <= last_reported_step:
|
|
return
|
|
accepted_step_callback(float(time))
|
|
last_reported_step = float(time)
|
|
|
|
def advance_to(
|
|
target_time: float, reported_terminal_time: float | None = None
|
|
) -> None:
|
|
nonlocal current_time, last_transition, same_time_transition_count, state
|
|
while current_time < target_time:
|
|
if cancel_check is not None and cancel_check():
|
|
raise _IntegrationCancelled
|
|
dt = min(config.max_step, target_time - current_time)
|
|
k1 = rhs(current_time, state)
|
|
k2 = rhs(
|
|
current_time + 0.5 * dt,
|
|
_vector_add(state, k1, 0.5 * dt),
|
|
)
|
|
k3 = rhs(
|
|
current_time + 0.5 * dt,
|
|
_vector_add(state, k2, 0.5 * dt),
|
|
)
|
|
k4 = rhs(current_time + dt, _vector_add(state, k3, dt))
|
|
next_state = [
|
|
value + (dt / 6.0) * (a + 2.0 * b + 2.0 * c + d)
|
|
for value, a, b, c, d in zip(state, k1, k2, k3, k4)
|
|
]
|
|
next_time = current_time + dt
|
|
|
|
transition: StateTransition | None = None
|
|
if state_transition_handler is not None:
|
|
step_start = current_time
|
|
step_state = list(state)
|
|
|
|
def dense_state(time: float) -> list[float]:
|
|
fraction = (float(time) - step_start) / (next_time - step_start)
|
|
return [
|
|
before + fraction * (after - before)
|
|
for before, after in zip(step_state, next_state)
|
|
]
|
|
|
|
candidate = state_transition_handler(
|
|
step_start,
|
|
list(step_state),
|
|
next_time,
|
|
list(next_state),
|
|
dense_state,
|
|
)
|
|
if candidate is not None:
|
|
requested_time = float(candidate.time)
|
|
candidate = _normalize_state_transition(
|
|
candidate,
|
|
step_start,
|
|
next_time,
|
|
len(state),
|
|
)
|
|
candidate = _align_transition_with_exact_endpoint(
|
|
candidate,
|
|
requested_time,
|
|
reported_terminal_time,
|
|
)
|
|
if not _is_repeated_state_transition(
|
|
candidate,
|
|
last_transition,
|
|
):
|
|
transition = candidate
|
|
|
|
if transition is not None:
|
|
same_time_transition_count = _next_same_time_transition_count(
|
|
transition,
|
|
last_transition,
|
|
same_time_transition_count,
|
|
)
|
|
current_time = transition.time
|
|
state = list(transition.state)
|
|
last_transition = transition
|
|
_append_or_replace_solution_sample(
|
|
times,
|
|
states,
|
|
current_time,
|
|
state,
|
|
)
|
|
else:
|
|
current_time = next_time
|
|
state = next_state
|
|
report_time = current_time
|
|
if (
|
|
reported_terminal_time is not None
|
|
and current_time >= target_time
|
|
):
|
|
report_time = reported_terminal_time
|
|
report_step(report_time)
|
|
|
|
try:
|
|
segment_ends = [*breakpoints, float(config.t_stop)]
|
|
for segment_index, segment_end in enumerate(segment_ends):
|
|
is_breakpoint = segment_index < len(breakpoints)
|
|
integration_end = (
|
|
math.nextafter(segment_end, -math.inf) if is_breakpoint else segment_end
|
|
)
|
|
|
|
while (
|
|
sample_index < len(sample_times)
|
|
and sample_times[sample_index] <= integration_end
|
|
):
|
|
sample_time = float(sample_times[sample_index])
|
|
advance_to(sample_time)
|
|
_append_solution_sample(times, states, sample_time, state)
|
|
sample_index += 1
|
|
|
|
advance_to(
|
|
integration_end,
|
|
segment_end if is_breakpoint else None,
|
|
)
|
|
|
|
if is_breakpoint:
|
|
current_time = float(segment_end)
|
|
report_step(current_time)
|
|
while (
|
|
sample_index < len(sample_times)
|
|
and sample_times[sample_index] <= segment_end
|
|
):
|
|
sample_time = float(sample_times[sample_index])
|
|
_append_solution_sample(times, states, sample_time, state)
|
|
sample_index += 1
|
|
except _IntegrationCancelled:
|
|
status = "cancelled"
|
|
message = "Simulation was stopped before reaching the requested end time."
|
|
_append_solution_sample(times, states, current_time, state)
|
|
except Exception as exc:
|
|
status = "failed"
|
|
message = str(exc)
|
|
error = exc
|
|
_append_solution_sample(times, states, current_time, state)
|
|
|
|
return ODESolution(
|
|
t=times,
|
|
y=states,
|
|
success=status == "completed",
|
|
message=message,
|
|
status=status,
|
|
error=error,
|
|
)
|
|
|
|
|
|
def _integrate_scipy_stepwise(
|
|
rhs: Callable[[float, list[float]], list[float]],
|
|
initial_state: list[float],
|
|
config: SolveIVPConfig,
|
|
t_eval: list[float] | None,
|
|
cancel_check: CancellationCheck,
|
|
accepted_step_callback: AcceptedStepCallback | None,
|
|
breakpoints: Sequence[float] = (),
|
|
state_transition_handler: StateTransitionHandler | None = None,
|
|
) -> ODESolution:
|
|
"""Initial stepwise integration path for breakpoints and state resets.
|
|
|
|
Recoverable physical-domain failures from rejected integrator trial states
|
|
restore the last accepted state and rebuild the same solver with a smaller
|
|
maximum/first step. Structural, algebraic, and ordinary model errors still
|
|
fail immediately.
|
|
"""
|
|
import numpy as np
|
|
from scipy.integrate import BDF, DOP853, LSODA, RK23, RK45, Radau
|
|
|
|
solver_types = {
|
|
"BDF": BDF,
|
|
"DOP853": DOP853,
|
|
"LSODA": LSODA,
|
|
"RK23": RK23,
|
|
"RK45": RK45,
|
|
"Radau": Radau,
|
|
}
|
|
solver_type = solver_types.get(config.method)
|
|
if solver_type is None:
|
|
raise ValueError(f"Unsupported integration method: {config.method}")
|
|
|
|
times = [float(config.t_start)]
|
|
states = [[float(value)] for value in initial_state]
|
|
last_accepted_time = float(config.t_start)
|
|
last_accepted_state = [float(value) for value in initial_state]
|
|
sample_times = [float(time) for time in (t_eval or [])]
|
|
sample_index = 0
|
|
while (
|
|
sample_index < len(sample_times)
|
|
and sample_times[sample_index] <= config.t_start
|
|
):
|
|
sample_index += 1
|
|
|
|
def cancellable_rhs(time, state):
|
|
if cancel_check():
|
|
raise _IntegrationCancelled
|
|
return rhs(float(time), [float(value) for value in state])
|
|
|
|
status: IntegrationStatus = "completed"
|
|
message = "The solver successfully reached the end of the integration interval."
|
|
error: Exception | None = None
|
|
last_transition: StateTransition | None = None
|
|
same_time_transition_count = 0
|
|
integration_progressed = False
|
|
last_reported_step: float | None = None
|
|
|
|
def cancellation_message() -> str:
|
|
return (
|
|
"Simulation was stopped before reaching the requested end time."
|
|
if integration_progressed
|
|
else "Simulation was stopped before integration started."
|
|
)
|
|
|
|
def report_step(time: float) -> None:
|
|
nonlocal last_reported_step
|
|
if accepted_step_callback is None:
|
|
return
|
|
if last_reported_step is not None and time <= last_reported_step:
|
|
return
|
|
accepted_step_callback(float(time))
|
|
last_reported_step = float(time)
|
|
|
|
segment_ends = [*breakpoints, float(config.t_stop)]
|
|
for segment_index, segment_end in enumerate(segment_ends):
|
|
if cancel_check():
|
|
status = "cancelled"
|
|
message = cancellation_message()
|
|
break
|
|
|
|
is_breakpoint = segment_index < len(breakpoints)
|
|
integration_end = (
|
|
math.nextafter(segment_end, -math.inf) if is_breakpoint else segment_end
|
|
)
|
|
has_integration_interval = integration_end > last_accepted_time
|
|
segment_max_step = float(config.max_step)
|
|
recoverable_retry_count = 0
|
|
last_recoverable_error: RecoverableTrialStateError | None = None
|
|
|
|
while has_integration_interval and last_accepted_time < integration_end:
|
|
if cancel_check():
|
|
status = "cancelled"
|
|
message = cancellation_message()
|
|
break
|
|
|
|
solver_options = {
|
|
"rtol": config.rtol,
|
|
"atol": config.atol,
|
|
"max_step": segment_max_step,
|
|
}
|
|
requested_first_step = (
|
|
0.1 * segment_max_step
|
|
if last_recoverable_error is not None
|
|
else config.first_step
|
|
)
|
|
if requested_first_step is not None:
|
|
solver_options["first_step"] = min(
|
|
requested_first_step,
|
|
integration_end - last_accepted_time,
|
|
)
|
|
|
|
try:
|
|
solver = solver_type(
|
|
cancellable_rhs,
|
|
last_accepted_time,
|
|
np.asarray(last_accepted_state, dtype=float),
|
|
integration_end,
|
|
**solver_options,
|
|
)
|
|
except _IntegrationCancelled:
|
|
status = "cancelled"
|
|
message = cancellation_message()
|
|
break
|
|
except RecoverableTrialStateError as exc:
|
|
recoverable_retry_count += 1
|
|
last_recoverable_error = exc
|
|
next_step = 0.5 * segment_max_step
|
|
minimum_step = 64.0 * math.ulp(max(abs(last_accepted_time), 1.0))
|
|
if recoverable_retry_count > 16 or next_step <= minimum_step:
|
|
status = "failed"
|
|
message = str(exc)
|
|
error = exc
|
|
break
|
|
segment_max_step = next_step
|
|
continue
|
|
except Exception as exc:
|
|
status = "failed"
|
|
message = str(exc)
|
|
error = exc
|
|
break
|
|
|
|
restart_at_transition = False
|
|
restart_after_recoverable = False
|
|
while solver.status == "running":
|
|
if cancel_check():
|
|
status = "cancelled"
|
|
message = (
|
|
"Simulation was stopped before reaching the requested end time."
|
|
)
|
|
break
|
|
|
|
step_start_time = last_accepted_time
|
|
step_start_state = list(last_accepted_state)
|
|
try:
|
|
step_message = solver.step()
|
|
except _IntegrationCancelled:
|
|
status = "cancelled"
|
|
message = (
|
|
"Simulation was stopped before reaching the requested end time."
|
|
)
|
|
break
|
|
except RecoverableTrialStateError as exc:
|
|
recoverable_retry_count += 1
|
|
last_recoverable_error = exc
|
|
attempted_step = segment_max_step
|
|
next_step = 0.5 * attempted_step
|
|
minimum_step = 64.0 * math.ulp(
|
|
max(abs(last_accepted_time), 1.0)
|
|
)
|
|
if recoverable_retry_count > 16 or next_step <= minimum_step:
|
|
status = "failed"
|
|
message = str(exc)
|
|
error = exc
|
|
break
|
|
segment_max_step = next_step
|
|
restart_after_recoverable = True
|
|
break
|
|
except Exception as exc:
|
|
status = "failed"
|
|
message = str(exc)
|
|
error = exc
|
|
break
|
|
|
|
integration_progressed = True
|
|
if solver.status == "failed":
|
|
if last_recoverable_error is not None:
|
|
recoverable_retry_count += 1
|
|
next_step = 0.5 * segment_max_step
|
|
minimum_step = 64.0 * math.ulp(
|
|
max(abs(last_accepted_time), 1.0)
|
|
)
|
|
if (
|
|
recoverable_retry_count <= 16
|
|
and next_step > minimum_step
|
|
):
|
|
segment_max_step = next_step
|
|
restart_after_recoverable = True
|
|
break
|
|
status = "failed"
|
|
message = str(step_message or "Integration step failed.")
|
|
break
|
|
|
|
step_end_time = float(solver.t)
|
|
step_end_state = [float(value) for value in solver.y]
|
|
dense_output = (
|
|
solver.dense_output()
|
|
if sample_times or state_transition_handler is not None
|
|
else None
|
|
)
|
|
|
|
transition: StateTransition | None = None
|
|
if state_transition_handler is not None:
|
|
assert dense_output is not None
|
|
|
|
def dense_state(time: float) -> list[float]:
|
|
return [float(value) for value in dense_output(float(time))]
|
|
|
|
try:
|
|
candidate = state_transition_handler(
|
|
step_start_time,
|
|
list(step_start_state),
|
|
step_end_time,
|
|
list(step_end_state),
|
|
dense_state,
|
|
)
|
|
if candidate is not None:
|
|
requested_time = float(candidate.time)
|
|
candidate = _normalize_state_transition(
|
|
candidate,
|
|
step_start_time,
|
|
step_end_time,
|
|
len(last_accepted_state),
|
|
)
|
|
candidate = _align_transition_with_exact_endpoint(
|
|
candidate,
|
|
requested_time,
|
|
float(segment_end) if is_breakpoint else None,
|
|
)
|
|
if not _is_repeated_state_transition(
|
|
candidate,
|
|
last_transition,
|
|
):
|
|
transition = candidate
|
|
except Exception as exc:
|
|
status = "failed"
|
|
message = str(exc)
|
|
error = exc
|
|
break
|
|
|
|
if transition is not None:
|
|
try:
|
|
same_time_transition_count = (
|
|
_next_same_time_transition_count(
|
|
transition,
|
|
last_transition,
|
|
same_time_transition_count,
|
|
)
|
|
)
|
|
except Exception as exc:
|
|
status = "failed"
|
|
message = str(exc)
|
|
error = exc
|
|
break
|
|
|
|
while (
|
|
sample_index < len(sample_times)
|
|
and sample_times[sample_index] < transition.time
|
|
):
|
|
sample_time = float(sample_times[sample_index])
|
|
assert dense_output is not None
|
|
sample_state = [
|
|
float(value) for value in dense_output(sample_time)
|
|
]
|
|
_append_solution_sample(
|
|
times,
|
|
states,
|
|
sample_time,
|
|
sample_state,
|
|
)
|
|
sample_index += 1
|
|
|
|
last_accepted_time = transition.time
|
|
last_accepted_state = list(transition.state)
|
|
last_transition = transition
|
|
_append_or_replace_solution_sample(
|
|
times,
|
|
states,
|
|
last_accepted_time,
|
|
last_accepted_state,
|
|
)
|
|
while (
|
|
sample_index < len(sample_times)
|
|
and sample_times[sample_index] <= last_accepted_time
|
|
):
|
|
sample_index += 1
|
|
report_step(last_accepted_time)
|
|
restart_at_transition = last_accepted_time < integration_end
|
|
break
|
|
|
|
last_accepted_time = step_end_time
|
|
last_accepted_state = step_end_state
|
|
recoverable_retry_count = 0
|
|
reported_time = (
|
|
float(segment_end)
|
|
if is_breakpoint and solver.status == "finished"
|
|
else last_accepted_time
|
|
)
|
|
if sample_times:
|
|
assert dense_output is not None
|
|
while (
|
|
sample_index < len(sample_times)
|
|
and sample_times[sample_index] <= last_accepted_time
|
|
):
|
|
sample_time = float(sample_times[sample_index])
|
|
sample_state = [
|
|
float(value) for value in dense_output(sample_time)
|
|
]
|
|
_append_solution_sample(
|
|
times,
|
|
states,
|
|
sample_time,
|
|
sample_state,
|
|
)
|
|
sample_index += 1
|
|
else:
|
|
_append_solution_sample(
|
|
times,
|
|
states,
|
|
reported_time,
|
|
last_accepted_state,
|
|
)
|
|
report_step(reported_time)
|
|
|
|
if status != "completed":
|
|
break
|
|
if restart_after_recoverable:
|
|
continue
|
|
if not restart_at_transition:
|
|
break
|
|
|
|
if status != "completed":
|
|
break
|
|
|
|
if is_breakpoint:
|
|
# The old equation is integrated only to the representable point just
|
|
# left of the event. The continuous state is then lifted to the exact
|
|
# event time, where the freshly constructed next solver sees the new
|
|
# equation immediately.
|
|
last_accepted_time = float(segment_end)
|
|
if sample_times:
|
|
while (
|
|
sample_index < len(sample_times)
|
|
and sample_times[sample_index] <= segment_end
|
|
):
|
|
sample_time = float(sample_times[sample_index])
|
|
_append_solution_sample(
|
|
times,
|
|
states,
|
|
sample_time,
|
|
last_accepted_state,
|
|
)
|
|
sample_index += 1
|
|
elif not has_integration_interval:
|
|
_append_solution_sample(
|
|
times,
|
|
states,
|
|
last_accepted_time,
|
|
last_accepted_state,
|
|
)
|
|
report_step(last_accepted_time)
|
|
|
|
if status != "completed":
|
|
_append_solution_sample(
|
|
times,
|
|
states,
|
|
last_accepted_time,
|
|
last_accepted_state,
|
|
)
|
|
|
|
return ODESolution(
|
|
t=times,
|
|
y=states,
|
|
success=status == "completed",
|
|
message=message,
|
|
status=status,
|
|
error=error,
|
|
)
|
|
|
|
|
|
def integrate_ode(
|
|
rhs: Callable[[float, list[float]], list[float]],
|
|
initial_state: list[float],
|
|
config: SolveIVPConfig,
|
|
t_eval: list[float] | None = None,
|
|
cancel_check: CancellationCheck | None = None,
|
|
accepted_step_callback: AcceptedStepCallback | None = None,
|
|
breakpoints: Sequence[float] | None = None,
|
|
state_transition_handler: StateTransitionHandler | None = None,
|
|
):
|
|
"""Integrate an ODE, optionally restarting at equation discontinuities.
|
|
|
|
Breakpoints are interpreted as right-continuous equation changes: the old
|
|
equation is integrated to the floating-point left limit, then a fresh solver
|
|
starts at the exact breakpoint with the unchanged continuous state.
|
|
|
|
A state transition handler inspects every accepted step using its dense
|
|
interpolant. When it returns a transition, samples before the event retain
|
|
the pre-event trajectory, the reset state is stored at the event, and a fresh
|
|
solver continues from that state.
|
|
"""
|
|
|
|
if (
|
|
state_transition_handler is not None
|
|
and config.t_stop < config.t_start
|
|
):
|
|
raise ValueError(
|
|
"State transition handling does not support reverse integration."
|
|
)
|
|
|
|
if config.t_stop == config.t_start:
|
|
return ODESolution(
|
|
t=[float(config.t_start)],
|
|
y=[[value] for value in initial_state],
|
|
success=True,
|
|
message="Skipped integration because t_start equals t_stop.",
|
|
)
|
|
|
|
normalized_breakpoints = _normalize_breakpoints(config, breakpoints)
|
|
|
|
try:
|
|
from scipy.integrate import solve_ivp
|
|
except ImportError:
|
|
if normalized_breakpoints:
|
|
return _runge_kutta_4_segmented(
|
|
rhs,
|
|
initial_state,
|
|
config,
|
|
t_eval,
|
|
normalized_breakpoints,
|
|
cancel_check,
|
|
accepted_step_callback,
|
|
state_transition_handler,
|
|
)
|
|
return _runge_kutta_4(
|
|
rhs,
|
|
initial_state,
|
|
config,
|
|
t_eval,
|
|
cancel_check,
|
|
accepted_step_callback,
|
|
state_transition_handler,
|
|
)
|
|
|
|
if (
|
|
cancel_check is not None
|
|
or normalized_breakpoints
|
|
or state_transition_handler is not None
|
|
):
|
|
return _integrate_scipy_stepwise(
|
|
rhs,
|
|
initial_state,
|
|
config,
|
|
t_eval,
|
|
cancel_check or (lambda: False),
|
|
accepted_step_callback,
|
|
normalized_breakpoints,
|
|
state_transition_handler,
|
|
)
|
|
|
|
solve_options = {
|
|
"fun": rhs,
|
|
"t_span": (config.t_start, config.t_stop),
|
|
"y0": initial_state,
|
|
"method": config.method,
|
|
"rtol": config.rtol,
|
|
"atol": config.atol,
|
|
"max_step": config.max_step,
|
|
"t_eval": t_eval,
|
|
}
|
|
if config.first_step is not None:
|
|
solve_options["first_step"] = config.first_step
|
|
return solve_ivp(**solve_options)
|