769 lines
27 KiB
Python
769 lines
27 KiB
Python
import math
|
|
import sys
|
|
import types
|
|
import unittest
|
|
from unittest.mock import patch
|
|
|
|
from app.simulation.core.errors import RecoverableTrialStateError
|
|
from app.simulation.solvers.solver import (
|
|
SolveIVPConfig,
|
|
StateTransition,
|
|
integrate_ode,
|
|
)
|
|
|
|
|
|
def _counting_fixed_step_solver(
|
|
step_size: float,
|
|
dense_output_times: list[float],
|
|
):
|
|
import numpy as np
|
|
|
|
class FixedStepSolver:
|
|
def __init__(self, _fun, t0, y0, t_bound, **_kwargs):
|
|
self.t = float(t0)
|
|
self.y = np.asarray(y0, dtype=float)
|
|
self.t_bound = float(t_bound)
|
|
self.status = "running"
|
|
self.nfev = 0
|
|
self.njev = 0
|
|
self.nlu = 0
|
|
|
|
def step(self):
|
|
self.t = min(self.t + step_size, self.t_bound)
|
|
self.y = np.asarray([self.t], dtype=float)
|
|
if self.t >= self.t_bound:
|
|
self.status = "finished"
|
|
return None
|
|
|
|
def dense_output(self):
|
|
dense_output_times.append(self.t)
|
|
return lambda time: np.asarray([float(time)], dtype=float)
|
|
|
|
return FixedStepSolver
|
|
|
|
|
|
class IntegrateOdeTests(unittest.TestCase):
|
|
def test_generic_solver_keeps_canonical_default_tolerance(self) -> None:
|
|
self.assertEqual(SolveIVPConfig().atol, 1.0e-8)
|
|
|
|
def test_scipy_solver_receives_step_size_controls(self) -> None:
|
|
calls: list[dict[str, object]] = []
|
|
|
|
def fake_solve_ivp(**kwargs):
|
|
calls.append(kwargs)
|
|
return object()
|
|
|
|
scipy_module = types.ModuleType("scipy")
|
|
integrate_module = types.ModuleType("scipy.integrate")
|
|
integrate_module.solve_ivp = fake_solve_ivp
|
|
scipy_module.integrate = integrate_module
|
|
|
|
with patch.dict(
|
|
sys.modules,
|
|
{"scipy": scipy_module, "scipy.integrate": integrate_module},
|
|
):
|
|
integrate_ode(
|
|
rhs=lambda _time, state: state,
|
|
initial_state=[1.0],
|
|
config=SolveIVPConfig(
|
|
t_start=0.0,
|
|
t_stop=1.0,
|
|
max_step=1.0e-4,
|
|
first_step=1.0e-8,
|
|
),
|
|
t_eval=[0.0, 1.0],
|
|
)
|
|
|
|
self.assertEqual(calls[0]["max_step"], 1.0e-4)
|
|
self.assertEqual(calls[0]["first_step"], 1.0e-8)
|
|
|
|
def test_scipy_solver_omits_unset_first_step(self) -> None:
|
|
calls: list[dict[str, object]] = []
|
|
|
|
def fake_solve_ivp(**kwargs):
|
|
calls.append(kwargs)
|
|
return object()
|
|
|
|
scipy_module = types.ModuleType("scipy")
|
|
integrate_module = types.ModuleType("scipy.integrate")
|
|
integrate_module.solve_ivp = fake_solve_ivp
|
|
scipy_module.integrate = integrate_module
|
|
|
|
with patch.dict(
|
|
sys.modules,
|
|
{"scipy": scipy_module, "scipy.integrate": integrate_module},
|
|
):
|
|
integrate_ode(
|
|
rhs=lambda _time, state: state,
|
|
initial_state=[1.0],
|
|
config=SolveIVPConfig(t_start=0.0, t_stop=1.0),
|
|
)
|
|
|
|
self.assertEqual(calls[0]["max_step"], 1.0e-3)
|
|
self.assertNotIn("first_step", calls[0])
|
|
|
|
def test_stepwise_solver_rebuilds_after_recoverable_trial_failure(self) -> None:
|
|
import numpy as np
|
|
import scipy.integrate
|
|
|
|
attempted_max_steps: list[float] = []
|
|
|
|
class RetryBdf:
|
|
def __init__(self, fun, t0, y0, t_bound, **kwargs):
|
|
self.fun = fun
|
|
self.t = float(t0)
|
|
self.y = np.asarray(y0, dtype=float)
|
|
self.t_bound = float(t_bound)
|
|
self.h_abs = float(kwargs["max_step"])
|
|
self.status = "running"
|
|
attempted_max_steps.append(self.h_abs)
|
|
|
|
def step(self):
|
|
if self.h_abs > 0.25:
|
|
raise RecoverableTrialStateError("trial state outside domain")
|
|
self.t = self.t_bound
|
|
self.status = "finished"
|
|
return None
|
|
|
|
def dense_output(self):
|
|
state = self.y.copy()
|
|
return lambda _time: state.copy()
|
|
|
|
with patch.object(scipy.integrate, "BDF", RetryBdf):
|
|
result = integrate_ode(
|
|
rhs=lambda _time, _state: [0.0],
|
|
initial_state=[1.0],
|
|
config=SolveIVPConfig(
|
|
t_start=0.0,
|
|
t_stop=1.0,
|
|
method="BDF",
|
|
max_step=1.0,
|
|
),
|
|
t_eval=[0.0, 1.0],
|
|
cancel_check=lambda: False,
|
|
)
|
|
|
|
self.assertTrue(result.success, result.message)
|
|
self.assertEqual(attempted_max_steps, [1.0, 0.5, 0.25])
|
|
self.assertEqual(result.t, [0.0, 1.0])
|
|
self.assertEqual(result.y, [[1.0, 1.0]])
|
|
|
|
def test_stepwise_solver_does_not_retry_ordinary_model_errors(self) -> None:
|
|
import numpy as np
|
|
import scipy.integrate
|
|
|
|
attempts = 0
|
|
|
|
class FailingBdf:
|
|
def __init__(self, fun, t0, y0, t_bound, **kwargs):
|
|
nonlocal attempts
|
|
attempts += 1
|
|
self.t = float(t0)
|
|
self.y = np.asarray(y0, dtype=float)
|
|
self.status = "running"
|
|
|
|
def step(self):
|
|
raise ValueError("structural model error")
|
|
|
|
with patch.object(scipy.integrate, "BDF", FailingBdf):
|
|
result = integrate_ode(
|
|
rhs=lambda _time, _state: [0.0],
|
|
initial_state=[1.0],
|
|
config=SolveIVPConfig(t_start=0.0, t_stop=1.0, method="BDF"),
|
|
cancel_check=lambda: False,
|
|
)
|
|
|
|
self.assertFalse(result.success)
|
|
self.assertEqual(result.message, "structural model error")
|
|
self.assertEqual(attempts, 1)
|
|
|
|
def test_scipy_stepwise_solver_can_cancel_before_start(self) -> None:
|
|
result = integrate_ode(
|
|
rhs=lambda _time, state: state,
|
|
initial_state=[1.0],
|
|
config=SolveIVPConfig(t_start=0.0, t_stop=1.0),
|
|
cancel_check=lambda: True,
|
|
)
|
|
|
|
self.assertFalse(result.success)
|
|
self.assertEqual(result.status, "cancelled")
|
|
self.assertEqual(result.t, [0.0])
|
|
|
|
def test_stepwise_solvers_build_dense_output_only_when_crossing_a_sample(
|
|
self,
|
|
) -> None:
|
|
import scipy.integrate
|
|
|
|
for method in ("BDF", "Radau", "RK45"):
|
|
with self.subTest(method=method):
|
|
dense_output_times: list[float] = []
|
|
fixed_step_solver = _counting_fixed_step_solver(
|
|
0.2,
|
|
dense_output_times,
|
|
)
|
|
|
|
with patch.object(scipy.integrate, method, fixed_step_solver):
|
|
result = integrate_ode(
|
|
rhs=lambda _time, _state: [1.0],
|
|
initial_state=[0.0],
|
|
config=SolveIVPConfig(
|
|
t_start=0.0,
|
|
t_stop=1.0,
|
|
method=method,
|
|
max_step=1.0,
|
|
),
|
|
t_eval=[0.0, 0.75, 1.0],
|
|
cancel_check=lambda: False,
|
|
)
|
|
|
|
self.assertTrue(result.success, result.message)
|
|
self.assertEqual(result.t, [0.0, 0.75, 1.0])
|
|
self.assertEqual(result.y, [[0.0, 0.75, 1.0]])
|
|
self.assertEqual(dense_output_times, [0.8, 1.0])
|
|
self.assertEqual(
|
|
result.solver_segments[0].accepted_step_count,
|
|
5,
|
|
)
|
|
|
|
def test_stepwise_solvers_keep_dense_output_for_state_event_detection(
|
|
self,
|
|
) -> None:
|
|
import scipy.integrate
|
|
|
|
for method in ("BDF", "Radau", "RK45"):
|
|
with self.subTest(method=method):
|
|
dense_output_times: list[float] = []
|
|
inspected_steps = 0
|
|
fixed_step_solver = _counting_fixed_step_solver(
|
|
0.25,
|
|
dense_output_times,
|
|
)
|
|
|
|
def inspect_state_event(
|
|
previous_time,
|
|
_previous_state,
|
|
current_time,
|
|
_current_state,
|
|
dense_state,
|
|
):
|
|
nonlocal inspected_steps
|
|
inspected_steps += 1
|
|
midpoint = 0.5 * (previous_time + current_time)
|
|
self.assertAlmostEqual(dense_state(midpoint)[0], midpoint)
|
|
return None
|
|
|
|
with patch.object(scipy.integrate, method, fixed_step_solver):
|
|
result = integrate_ode(
|
|
rhs=lambda _time, _state: [1.0],
|
|
initial_state=[0.0],
|
|
config=SolveIVPConfig(
|
|
t_start=0.0,
|
|
t_stop=1.0,
|
|
method=method,
|
|
max_step=1.0,
|
|
),
|
|
t_eval=[0.0, 1.0],
|
|
state_transition_handler=inspect_state_event,
|
|
)
|
|
|
|
self.assertTrue(result.success, result.message)
|
|
self.assertEqual(len(dense_output_times), 4)
|
|
self.assertEqual(inspected_steps, 4)
|
|
|
|
def test_stepwise_dense_output_preserves_adjacent_float_samples(self) -> None:
|
|
import scipy.integrate
|
|
|
|
dense_output_times: list[float] = []
|
|
adjacent_time = math.nextafter(0.5, math.inf)
|
|
half_interval_solver = _counting_fixed_step_solver(
|
|
0.5,
|
|
dense_output_times,
|
|
)
|
|
|
|
with patch.object(scipy.integrate, "RK45", half_interval_solver):
|
|
result = integrate_ode(
|
|
rhs=lambda _time, _state: [1.0],
|
|
initial_state=[0.0],
|
|
config=SolveIVPConfig(
|
|
t_start=0.0,
|
|
t_stop=1.0,
|
|
method="RK45",
|
|
max_step=1.0,
|
|
),
|
|
t_eval=[0.0, 0.5, adjacent_time, 1.0],
|
|
cancel_check=lambda: False,
|
|
)
|
|
|
|
self.assertTrue(result.success, result.message)
|
|
self.assertEqual(result.t, [0.0, 0.5, adjacent_time, 1.0])
|
|
self.assertEqual(result.y[0], result.t)
|
|
self.assertEqual(dense_output_times, [0.5, 1.0])
|
|
|
|
def test_dense_output_pruning_preserves_cancelled_partial_samples(self) -> None:
|
|
import scipy.integrate
|
|
|
|
cancellation_requested = False
|
|
dense_output_times: list[float] = []
|
|
fixed_step_solver = _counting_fixed_step_solver(
|
|
0.2,
|
|
dense_output_times,
|
|
)
|
|
|
|
def request_cancel(time: float) -> None:
|
|
nonlocal cancellation_requested
|
|
cancellation_requested = time >= 0.4
|
|
|
|
with patch.object(scipy.integrate, "BDF", fixed_step_solver):
|
|
result = integrate_ode(
|
|
rhs=lambda _time, _state: [1.0],
|
|
initial_state=[0.0],
|
|
config=SolveIVPConfig(
|
|
t_start=0.0,
|
|
t_stop=1.0,
|
|
method="BDF",
|
|
max_step=1.0,
|
|
),
|
|
t_eval=[0.0, 0.3, 0.8, 1.0],
|
|
cancel_check=lambda: cancellation_requested,
|
|
accepted_step_callback=request_cancel,
|
|
)
|
|
|
|
self.assertFalse(result.success)
|
|
self.assertEqual(result.status, "cancelled")
|
|
self.assertEqual(result.t, [0.0, 0.3, 0.4])
|
|
self.assertEqual(result.y, [[0.0, 0.3, 0.4]])
|
|
self.assertEqual(dense_output_times, [0.4])
|
|
|
|
def test_segmented_bdf_uses_left_limit_and_restarts_at_event(self) -> None:
|
|
import scipy.integrate
|
|
|
|
event_time = 0.5
|
|
actual_bdf = scipy.integrate.BDF
|
|
starts: list[float] = []
|
|
bounds: list[float] = []
|
|
call_times: list[list[float]] = []
|
|
|
|
class RecordingBDF(actual_bdf):
|
|
def __init__(self, fun, t0, y0, t_bound, **kwargs):
|
|
starts.append(float(t0))
|
|
bounds.append(float(t_bound))
|
|
segment_calls: list[float] = []
|
|
call_times.append(segment_calls)
|
|
|
|
def recording_fun(time, state):
|
|
segment_calls.append(float(time))
|
|
return fun(time, state)
|
|
|
|
super().__init__(recording_fun, t0, y0, t_bound, **kwargs)
|
|
|
|
with patch.object(scipy.integrate, "BDF", RecordingBDF):
|
|
result = integrate_ode(
|
|
rhs=lambda time, _state: [1.0 if time < event_time else 2.0],
|
|
initial_state=[0.0],
|
|
config=SolveIVPConfig(
|
|
t_start=0.0,
|
|
t_stop=1.0,
|
|
method="BDF",
|
|
max_step=0.1,
|
|
first_step=0.8,
|
|
),
|
|
t_eval=[0.0, event_time, event_time, 1.0],
|
|
breakpoints=[event_time],
|
|
)
|
|
|
|
self.assertTrue(result.success, result.message)
|
|
self.assertEqual(starts, [0.0, event_time])
|
|
self.assertEqual(bounds[0], math.nextafter(event_time, -math.inf))
|
|
self.assertEqual(bounds[1], 1.0)
|
|
self.assertTrue(call_times[0])
|
|
self.assertTrue(all(time < event_time for time in call_times[0]))
|
|
self.assertTrue(any(time >= event_time for time in call_times[1]))
|
|
self.assertEqual(result.t, [0.0, event_time, 1.0])
|
|
self.assertAlmostEqual(result.y[0][-1], 1.5, places=5)
|
|
|
|
def test_segmented_implicit_solvers_merge_samples_and_report_progress(self) -> None:
|
|
event_time = 0.4
|
|
|
|
for method in ("BDF", "Radau"):
|
|
with self.subTest(method=method):
|
|
callback_times: list[float] = []
|
|
result = integrate_ode(
|
|
rhs=lambda time, _state: [1.0 if time < event_time else 3.0],
|
|
initial_state=[0.0],
|
|
config=SolveIVPConfig(
|
|
t_start=0.0,
|
|
t_stop=1.0,
|
|
method=method,
|
|
max_step=0.05,
|
|
first_step=0.9,
|
|
),
|
|
t_eval=[0.0, event_time, event_time, 0.7, 1.0],
|
|
accepted_step_callback=callback_times.append,
|
|
breakpoints=[event_time, event_time],
|
|
)
|
|
|
|
self.assertTrue(result.success, result.message)
|
|
self.assertEqual(result.t, [0.0, event_time, 0.7, 1.0])
|
|
self.assertAlmostEqual(result.y[0][-1], 2.2, places=5)
|
|
self.assertEqual(callback_times.count(event_time), 1)
|
|
self.assertTrue(
|
|
all(
|
|
earlier < later
|
|
for earlier, later in zip(
|
|
callback_times,
|
|
callback_times[1:],
|
|
)
|
|
)
|
|
)
|
|
|
|
def test_segmented_solver_reports_implicit_work_by_event_segment(self) -> None:
|
|
import numpy as np
|
|
import scipy.integrate
|
|
|
|
class CountingBDF:
|
|
def __init__(self, _fun, t0, y0, t_bound, **_kwargs):
|
|
self.t = float(t0)
|
|
self.y = np.asarray(y0, dtype=float)
|
|
self.t_bound = float(t_bound)
|
|
self.status = "running"
|
|
self.nfev = 2
|
|
self.njev = 1
|
|
self.nlu = 0
|
|
|
|
def step(self):
|
|
self.t = self.t_bound
|
|
self.nfev += 3
|
|
self.nlu += 2
|
|
self.status = "finished"
|
|
return None
|
|
|
|
def dense_output(self):
|
|
state = self.y.copy()
|
|
return lambda _time: state.copy()
|
|
|
|
with patch.object(scipy.integrate, "BDF", CountingBDF):
|
|
result = integrate_ode(
|
|
rhs=lambda _time, _state: [0.0],
|
|
initial_state=[1.0],
|
|
config=SolveIVPConfig(
|
|
t_start=0.0,
|
|
t_stop=1.0,
|
|
method="BDF",
|
|
max_step=1.0,
|
|
),
|
|
t_eval=[0.0, 0.4, 1.0],
|
|
breakpoints=[0.4],
|
|
)
|
|
|
|
self.assertTrue(result.success, result.message)
|
|
self.assertEqual(len(result.solver_segments), 2)
|
|
self.assertEqual(
|
|
[segment.as_dict() for segment in result.solver_segments],
|
|
[
|
|
{
|
|
"startTime": 0.0,
|
|
"requestedStopTime": 0.4,
|
|
"simulatedUntil": 0.4,
|
|
"nfev": 5,
|
|
"njev": 1,
|
|
"nlu": 2,
|
|
"acceptedStepCount": 1,
|
|
"solverStartCount": 1,
|
|
"stateTransitionCount": 0,
|
|
"recoverableRetryCount": 0,
|
|
},
|
|
{
|
|
"startTime": 0.4,
|
|
"requestedStopTime": 1.0,
|
|
"simulatedUntil": 1.0,
|
|
"nfev": 5,
|
|
"njev": 1,
|
|
"nlu": 2,
|
|
"acceptedStepCount": 1,
|
|
"solverStartCount": 1,
|
|
"stateTransitionCount": 0,
|
|
"recoverableRetryCount": 0,
|
|
},
|
|
],
|
|
)
|
|
|
|
def test_segmented_solver_can_cancel_after_crossing_a_breakpoint(self) -> None:
|
|
callback_times: list[float] = []
|
|
cancellation_requested = False
|
|
|
|
def record_progress(time: float) -> None:
|
|
nonlocal cancellation_requested
|
|
callback_times.append(time)
|
|
cancellation_requested = time >= 0.55
|
|
|
|
result = integrate_ode(
|
|
rhs=lambda _time, _state: [1.0],
|
|
initial_state=[0.0],
|
|
config=SolveIVPConfig(
|
|
t_start=0.0,
|
|
t_stop=1.0,
|
|
method="BDF",
|
|
max_step=0.05,
|
|
),
|
|
t_eval=[0.0, 0.2, 0.4, 0.6, 0.8, 1.0],
|
|
cancel_check=lambda: cancellation_requested,
|
|
accepted_step_callback=record_progress,
|
|
breakpoints=[0.4, 0.8],
|
|
)
|
|
|
|
self.assertFalse(result.success)
|
|
self.assertEqual(result.status, "cancelled")
|
|
self.assertIn(0.4, callback_times)
|
|
self.assertGreater(callback_times[-1], 0.4)
|
|
self.assertTrue(
|
|
all(
|
|
earlier < later
|
|
for earlier, later in zip(callback_times, callback_times[1:])
|
|
)
|
|
)
|
|
self.assertEqual(result.t, sorted(set(result.t)))
|
|
|
|
def test_state_transition_resets_at_root_and_discards_step_overshoot(self) -> None:
|
|
event_time = 0.35
|
|
event_enabled = True
|
|
|
|
def transition_handler(
|
|
previous_time,
|
|
previous_state,
|
|
current_time,
|
|
current_state,
|
|
dense_state,
|
|
):
|
|
nonlocal event_enabled
|
|
if (
|
|
not event_enabled
|
|
or previous_state[0] >= event_time
|
|
or current_state[0] < event_time
|
|
):
|
|
return None
|
|
lower = previous_time
|
|
upper = current_time
|
|
for _iteration in range(60):
|
|
middle = 0.5 * (lower + upper)
|
|
if dense_state(middle)[0] >= event_time:
|
|
upper = middle
|
|
else:
|
|
lower = middle
|
|
event_enabled = False
|
|
return StateTransition(time=upper, state=[0.0])
|
|
|
|
result = integrate_ode(
|
|
rhs=lambda _time, _state: [1.0],
|
|
initial_state=[0.0],
|
|
config=SolveIVPConfig(
|
|
t_start=0.0,
|
|
t_stop=1.0,
|
|
method="BDF",
|
|
max_step=0.5,
|
|
),
|
|
t_eval=[0.0, event_time, 0.4, 1.0],
|
|
state_transition_handler=transition_handler,
|
|
)
|
|
|
|
self.assertTrue(result.success, result.message)
|
|
self.assertEqual(result.t, [0.0, event_time, 0.4, 1.0])
|
|
self.assertAlmostEqual(result.y[0][1], 0.0, places=12)
|
|
self.assertAlmostEqual(result.y[0][2], 0.05, places=8)
|
|
self.assertAlmostEqual(result.y[0][-1], 0.65, places=8)
|
|
|
|
def test_state_transitions_chain_at_same_time_until_state_repeats(self) -> None:
|
|
event_time = 0.25
|
|
stage = 0
|
|
returned_reset_states: list[float] = []
|
|
|
|
def transition_handler(
|
|
previous_time,
|
|
_previous_state,
|
|
current_time,
|
|
_current_state,
|
|
_dense_state,
|
|
):
|
|
nonlocal stage
|
|
if stage == 0 and previous_time <= event_time <= current_time:
|
|
stage = 1
|
|
returned_reset_states.append(10.0)
|
|
return StateTransition(time=event_time, state=[10.0])
|
|
if stage == 1 and previous_time == event_time:
|
|
stage = 2
|
|
returned_reset_states.append(20.0)
|
|
return StateTransition(time=event_time, state=[20.0])
|
|
if stage == 2 and previous_time == event_time:
|
|
returned_reset_states.append(20.0)
|
|
return StateTransition(time=event_time, state=[20.0])
|
|
return None
|
|
|
|
result = integrate_ode(
|
|
rhs=lambda _time, _state: [1.0],
|
|
initial_state=[0.0],
|
|
config=SolveIVPConfig(
|
|
t_start=0.0,
|
|
t_stop=1.0,
|
|
method="BDF",
|
|
max_step=0.4,
|
|
),
|
|
t_eval=[0.0, event_time, 1.0],
|
|
state_transition_handler=transition_handler,
|
|
)
|
|
|
|
self.assertTrue(result.success, result.message)
|
|
self.assertEqual(returned_reset_states, [10.0, 20.0, 20.0])
|
|
self.assertEqual(result.t, [0.0, event_time, 1.0])
|
|
self.assertEqual(result.y[0][1], 20.0)
|
|
self.assertAlmostEqual(result.y[0][-1], 20.75, places=8)
|
|
|
|
def test_state_transition_chain_has_a_finite_guard(self) -> None:
|
|
event_time = 0.25
|
|
reset_count = 0
|
|
|
|
def transition_handler(
|
|
previous_time,
|
|
_previous_state,
|
|
current_time,
|
|
_current_state,
|
|
_dense_state,
|
|
):
|
|
nonlocal reset_count
|
|
if previous_time <= event_time <= current_time:
|
|
reset_count += 1
|
|
return StateTransition(
|
|
time=event_time,
|
|
state=[float(reset_count)],
|
|
)
|
|
return None
|
|
|
|
result = integrate_ode(
|
|
rhs=lambda _time, _state: [1.0],
|
|
initial_state=[0.0],
|
|
config=SolveIVPConfig(
|
|
t_start=0.0,
|
|
t_stop=1.0,
|
|
method="BDF",
|
|
max_step=0.4,
|
|
),
|
|
state_transition_handler=transition_handler,
|
|
)
|
|
|
|
self.assertFalse(result.success)
|
|
self.assertEqual(result.status, "failed")
|
|
self.assertIn("64 chained resets", result.message)
|
|
|
|
def test_stepwise_solver_preserves_adjacent_float_samples(self) -> None:
|
|
adjacent_time = math.nextafter(0.5, math.inf)
|
|
|
|
result = integrate_ode(
|
|
rhs=lambda _time, _state: [1.0],
|
|
initial_state=[0.0],
|
|
config=SolveIVPConfig(
|
|
t_start=0.0,
|
|
t_stop=1.0,
|
|
method="RK45",
|
|
max_step=0.4,
|
|
),
|
|
t_eval=[0.0, 0.5, adjacent_time, 1.0],
|
|
state_transition_handler=lambda *_args: None,
|
|
)
|
|
|
|
self.assertTrue(result.success, result.message)
|
|
self.assertEqual(result.t, [0.0, 0.5, adjacent_time, 1.0])
|
|
|
|
def test_state_transition_at_breakpoint_uses_exact_breakpoint_sample(self) -> None:
|
|
event_time = 0.5
|
|
integration_left_limit = math.nextafter(event_time, -math.inf)
|
|
event_enabled = True
|
|
|
|
def transition_handler(
|
|
previous_time,
|
|
_previous_state,
|
|
current_time,
|
|
_current_state,
|
|
_dense_state,
|
|
):
|
|
nonlocal event_enabled
|
|
if (
|
|
event_enabled
|
|
and previous_time <= integration_left_limit <= current_time
|
|
):
|
|
event_enabled = False
|
|
return StateTransition(time=event_time, state=[7.0])
|
|
return None
|
|
|
|
result = integrate_ode(
|
|
rhs=lambda _time, _state: [1.0],
|
|
initial_state=[0.0],
|
|
config=SolveIVPConfig(
|
|
t_start=0.0,
|
|
t_stop=1.0,
|
|
method="BDF",
|
|
max_step=0.2,
|
|
),
|
|
t_eval=[0.0, event_time, 1.0],
|
|
breakpoints=[event_time],
|
|
state_transition_handler=transition_handler,
|
|
)
|
|
|
|
self.assertTrue(result.success, result.message)
|
|
self.assertEqual(result.t, [0.0, event_time, 1.0])
|
|
self.assertEqual(result.y[0][1], 7.0)
|
|
self.assertAlmostEqual(result.y[0][-1], 7.5, places=8)
|
|
|
|
def test_cancellation_after_state_transition_reports_partial_progress(self) -> None:
|
|
event_time = 0.25
|
|
cancellation_requested = False
|
|
event_enabled = True
|
|
|
|
def transition_handler(
|
|
previous_time,
|
|
_previous_state,
|
|
current_time,
|
|
_current_state,
|
|
_dense_state,
|
|
):
|
|
nonlocal cancellation_requested, event_enabled
|
|
if event_enabled and previous_time <= event_time <= current_time:
|
|
event_enabled = False
|
|
cancellation_requested = True
|
|
return StateTransition(time=event_time, state=[0.0])
|
|
return None
|
|
|
|
result = integrate_ode(
|
|
rhs=lambda _time, _state: [1.0],
|
|
initial_state=[0.0],
|
|
config=SolveIVPConfig(
|
|
t_start=0.0,
|
|
t_stop=1.0,
|
|
method="BDF",
|
|
max_step=0.4,
|
|
),
|
|
cancel_check=lambda: cancellation_requested,
|
|
state_transition_handler=transition_handler,
|
|
)
|
|
|
|
self.assertFalse(result.success)
|
|
self.assertEqual(result.status, "cancelled")
|
|
self.assertEqual(
|
|
result.message,
|
|
"Simulation was stopped before reaching the requested end time.",
|
|
)
|
|
self.assertEqual(result.t[-1], event_time)
|
|
|
|
def test_state_transition_handler_rejects_reverse_integration(self) -> None:
|
|
with self.assertRaisesRegex(
|
|
ValueError,
|
|
"does not support reverse integration",
|
|
):
|
|
integrate_ode(
|
|
rhs=lambda _time, _state: [1.0],
|
|
initial_state=[0.0],
|
|
config=SolveIVPConfig(t_start=1.0, t_stop=0.0),
|
|
state_transition_handler=lambda *_args: None,
|
|
)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|