203 lines
7.2 KiB
Python
203 lines
7.2 KiB
Python
import math
|
|
import sys
|
|
import types
|
|
import unittest
|
|
from unittest.mock import patch
|
|
|
|
from app.simulation.solvers.solver import SolveIVPConfig, integrate_ode
|
|
|
|
|
|
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_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_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_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)))
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|