1418 lines
50 KiB
Python
1418 lines
50 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 (
|
|
IntegrationCancelled,
|
|
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_stepwise_jacobian_cancellation_returns_cancelled_solution(self) -> None:
|
|
class CancellingJacobian:
|
|
def start_segment(self) -> None:
|
|
pass
|
|
|
|
def observe(self, *_args) -> None:
|
|
pass
|
|
|
|
def __call__(self, _time, _state):
|
|
raise IntegrationCancelled
|
|
|
|
result = integrate_ode(
|
|
rhs=lambda _time, state: [-float(state[0])],
|
|
initial_state=[1.0],
|
|
config=SolveIVPConfig(t_stop=1.0, method="BDF"),
|
|
cancel_check=lambda: False,
|
|
jac=CancellingJacobian(),
|
|
)
|
|
|
|
self.assertFalse(result.success)
|
|
self.assertEqual(result.status, "cancelled")
|
|
self.assertEqual(result.t, [0.0])
|
|
|
|
def test_stepwise_explicit_solver_ignores_callable_jacobian(self) -> None:
|
|
class UnexpectedJacobian:
|
|
call_count = 0
|
|
|
|
def start_segment(self) -> None:
|
|
self.call_count += 1
|
|
|
|
def observe(self, *_args) -> None:
|
|
self.call_count += 1
|
|
|
|
def __call__(self, _time, _state):
|
|
self.call_count += 1
|
|
raise AssertionError("Explicit solver evaluated Jacobian")
|
|
|
|
jacobian = UnexpectedJacobian()
|
|
result = integrate_ode(
|
|
rhs=lambda _time, _state: [1.0],
|
|
initial_state=[0.0],
|
|
config=SolveIVPConfig(t_stop=0.01, method="RK45"),
|
|
state_transition_handler=lambda *_args: None,
|
|
jac=jacobian,
|
|
)
|
|
|
|
self.assertTrue(result.success, result.message)
|
|
self.assertEqual(jacobian.call_count, 0)
|
|
|
|
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_direct_implicit_solver_receives_callable_jacobian(self) -> None:
|
|
calls: list[dict[str, object]] = []
|
|
observed: list[tuple[float, list[float], list[float]]] = []
|
|
|
|
class RecordingJacobian:
|
|
def __init__(self) -> None:
|
|
self.segment_count = 0
|
|
|
|
def start_segment(self) -> None:
|
|
self.segment_count += 1
|
|
|
|
def observe(self, time, state, derivative) -> None:
|
|
observed.append(
|
|
(
|
|
float(time),
|
|
[float(value) for value in state],
|
|
[float(value) for value in derivative],
|
|
)
|
|
)
|
|
|
|
def __call__(self, _time, _state):
|
|
return [[1.0]]
|
|
|
|
def fake_solve_ivp(**kwargs):
|
|
calls.append(kwargs)
|
|
kwargs["fun"](0.0, [2.0])
|
|
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
|
|
jacobian = RecordingJacobian()
|
|
|
|
with patch.dict(
|
|
sys.modules,
|
|
{"scipy": scipy_module, "scipy.integrate": integrate_module},
|
|
):
|
|
integrate_ode(
|
|
rhs=lambda _time, state: [2.0 * state[0]],
|
|
initial_state=[1.0],
|
|
config=SolveIVPConfig(t_start=0.0, t_stop=1.0, method="BDF"),
|
|
jac_sparsity=[[True]],
|
|
jac=jacobian,
|
|
)
|
|
|
|
self.assertIs(calls[0]["jac"], jacobian)
|
|
self.assertNotIn("jac_sparsity", calls[0])
|
|
self.assertEqual(jacobian.segment_count, 1)
|
|
self.assertEqual(observed, [(0.0, [2.0], [4.0])])
|
|
|
|
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_recoverable_trial_retries_default_keeps_direct_scipy_route(
|
|
self,
|
|
) -> None:
|
|
import scipy.integrate
|
|
|
|
expected = object()
|
|
with patch.object(
|
|
scipy.integrate,
|
|
"solve_ivp",
|
|
return_value=expected,
|
|
) as direct_solve:
|
|
result = integrate_ode(
|
|
rhs=lambda _time, state: state,
|
|
initial_state=[1.0],
|
|
config=SolveIVPConfig(t_start=0.0, t_stop=1.0),
|
|
)
|
|
|
|
self.assertIs(result, expected)
|
|
direct_solve.assert_called_once()
|
|
|
|
def test_recoverable_stepwise_route_matches_direct_scipy_without_failures(
|
|
self,
|
|
) -> None:
|
|
import numpy as np
|
|
|
|
sample_times = [0.0, 0.025, 0.05, 0.075, 0.1]
|
|
|
|
def rhs(time, state):
|
|
return [-2.0 * float(state[0]) + math.sin(float(time))]
|
|
|
|
for method in ("RK45", "BDF"):
|
|
with self.subTest(method=method):
|
|
config = SolveIVPConfig(
|
|
t_start=0.0,
|
|
t_stop=0.1,
|
|
method=method,
|
|
rtol=1.0e-9,
|
|
atol=1.0e-12,
|
|
max_step=0.01,
|
|
)
|
|
direct = integrate_ode(
|
|
rhs=rhs,
|
|
initial_state=[1.0],
|
|
config=config,
|
|
t_eval=sample_times,
|
|
)
|
|
stepwise = integrate_ode(
|
|
rhs=rhs,
|
|
initial_state=[1.0],
|
|
config=config,
|
|
t_eval=sample_times,
|
|
recoverable_trial_retries=True,
|
|
)
|
|
|
|
self.assertTrue(direct.success, direct.message)
|
|
self.assertTrue(stepwise.success, stepwise.message)
|
|
np.testing.assert_array_equal(direct.t, stepwise.t)
|
|
np.testing.assert_allclose(
|
|
direct.y,
|
|
stepwise.y,
|
|
rtol=1.0e-12,
|
|
atol=1.0e-14,
|
|
)
|
|
self.assertEqual(
|
|
int(direct.nfev),
|
|
sum(segment.nfev for segment in stepwise.solver_segments),
|
|
)
|
|
self.assertEqual(
|
|
int(getattr(direct, "njev", 0)),
|
|
sum(segment.njev for segment in stepwise.solver_segments),
|
|
)
|
|
self.assertEqual(
|
|
int(getattr(direct, "nlu", 0)),
|
|
sum(segment.nlu for segment in stepwise.solver_segments),
|
|
)
|
|
|
|
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],
|
|
recoverable_trial_retries=True,
|
|
)
|
|
|
|
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_retries_recoverable_constructor_failure(self) -> None:
|
|
import numpy as np
|
|
import scipy.integrate
|
|
|
|
attempts: list[dict[str, object]] = []
|
|
|
|
class ConstructorRetryBdf:
|
|
def __init__(self, _fun, t0, y0, t_bound, **kwargs):
|
|
attempts.append(dict(kwargs))
|
|
if float(kwargs["max_step"]) > 0.25:
|
|
raise RecoverableTrialStateError("constructor trial failed")
|
|
self.t = float(t0)
|
|
self.y = np.asarray(y0, dtype=float)
|
|
self.t_bound = float(t_bound)
|
|
self.status = "running"
|
|
|
|
def step(self):
|
|
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", ConstructorRetryBdf):
|
|
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,
|
|
),
|
|
cancel_check=lambda: False,
|
|
)
|
|
|
|
self.assertTrue(result.success, result.message)
|
|
self.assertNotIn("first_step", attempts[0])
|
|
self.assertEqual(
|
|
[attempt["max_step"] for attempt in attempts],
|
|
[1.0, 0.5, 0.25],
|
|
)
|
|
self.assertEqual(
|
|
[attempt.get("first_step") for attempt in attempts],
|
|
[None, 0.5, 0.25],
|
|
)
|
|
segment = result.solver_segments[0]
|
|
self.assertEqual(segment.recoverable_retry_count, 2)
|
|
self.assertEqual(
|
|
[retry.phase for retry in segment.recoverable_retries],
|
|
["constructor", "constructor"],
|
|
)
|
|
self.assertEqual(
|
|
segment.as_dict()["recoverableRetries"][0],
|
|
{
|
|
"phase": "constructor",
|
|
"attemptedStep": 1.0,
|
|
"reason": "constructor trial failed",
|
|
"nextMaxStep": 0.5,
|
|
"nextFirstStep": 0.5,
|
|
},
|
|
)
|
|
|
|
def test_stepwise_retry_uses_solver_actual_step_not_segment_maximum(
|
|
self,
|
|
) -> None:
|
|
import numpy as np
|
|
import scipy.integrate
|
|
|
|
attempts: list[dict[str, object]] = []
|
|
|
|
class ActualStepRetryBdf:
|
|
def __init__(self, _fun, t0, y0, t_bound, **kwargs):
|
|
self.attempt_index = len(attempts)
|
|
attempts.append(dict(kwargs))
|
|
self.t = float(t0)
|
|
self.y = np.asarray(y0, dtype=float)
|
|
self.t_bound = float(t_bound)
|
|
self.h_abs = min(0.04, float(kwargs["max_step"]))
|
|
self.step_size = min(0.03, float(kwargs["max_step"]))
|
|
self.status = "running"
|
|
|
|
def step(self):
|
|
if self.attempt_index == 0:
|
|
raise RecoverableTrialStateError("small actual trial failed")
|
|
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", ActualStepRetryBdf):
|
|
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,
|
|
),
|
|
cancel_check=lambda: False,
|
|
)
|
|
|
|
self.assertTrue(result.success, result.message)
|
|
self.assertNotIn("first_step", attempts[0])
|
|
self.assertEqual(attempts[1]["max_step"], 0.02)
|
|
self.assertEqual(attempts[1]["first_step"], 0.02)
|
|
retry = result.solver_segments[0].recoverable_retries[0]
|
|
self.assertEqual(retry.phase, "step")
|
|
self.assertEqual(retry.attempted_step, 0.04)
|
|
self.assertEqual(retry.next_max_step, 0.02)
|
|
|
|
def test_stepwise_retry_falls_back_to_previous_step_size_without_h_abs(
|
|
self,
|
|
) -> None:
|
|
import numpy as np
|
|
import scipy.integrate
|
|
|
|
for invalid_h_abs in (None, math.nan):
|
|
with self.subTest(h_abs=invalid_h_abs):
|
|
attempts: list[dict[str, object]] = []
|
|
|
|
class StepSizeFallbackBdf:
|
|
def __init__(self, _fun, t0, y0, t_bound, **kwargs):
|
|
self.attempt_index = len(attempts)
|
|
attempts.append(dict(kwargs))
|
|
self.t = float(t0)
|
|
self.y = np.asarray(y0, dtype=float)
|
|
self.t_bound = float(t_bound)
|
|
self.h_abs = invalid_h_abs
|
|
self.step_size = 0.06
|
|
self.status = "running"
|
|
|
|
def step(self):
|
|
if self.attempt_index == 0:
|
|
raise RecoverableTrialStateError(
|
|
"trial without valid h_abs"
|
|
)
|
|
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",
|
|
StepSizeFallbackBdf,
|
|
):
|
|
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,
|
|
),
|
|
cancel_check=lambda: False,
|
|
)
|
|
|
|
self.assertTrue(result.success, result.message)
|
|
self.assertEqual(attempts[1]["max_step"], 0.03)
|
|
self.assertEqual(attempts[1]["first_step"], 0.03)
|
|
retry = result.solver_segments[0].recoverable_retries[0]
|
|
self.assertEqual(retry.attempted_step, 0.06)
|
|
self.assertEqual(retry.next_max_step, 0.03)
|
|
|
|
def test_accepted_step_clears_recoverable_retry_state(self) -> None:
|
|
import numpy as np
|
|
import scipy.integrate
|
|
|
|
attempts: list[dict[str, object]] = []
|
|
caps_after_accepted_retry: list[float] = []
|
|
|
|
class FailureAfterAcceptedRetryBdf:
|
|
def __init__(self, _fun, t0, y0, t_bound, **kwargs):
|
|
self.attempt_index = len(attempts)
|
|
attempts.append(dict(kwargs))
|
|
self.t = float(t0)
|
|
self.y = np.asarray(y0, dtype=float)
|
|
self.t_bound = float(t_bound)
|
|
self.h_abs = min(0.2, float(kwargs["max_step"]))
|
|
self.status = "running"
|
|
self.step_count = 0
|
|
|
|
def step(self):
|
|
self.step_count += 1
|
|
if self.attempt_index == 0:
|
|
raise RecoverableTrialStateError("retry once")
|
|
if self.step_count == 1:
|
|
self.t = 0.25
|
|
return None
|
|
caps_after_accepted_retry.append(float(self.max_step))
|
|
self.status = "failed"
|
|
return "ordinary failure after accepted step"
|
|
|
|
with patch.object(
|
|
scipy.integrate,
|
|
"BDF",
|
|
FailureAfterAcceptedRetryBdf,
|
|
):
|
|
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,
|
|
),
|
|
cancel_check=lambda: False,
|
|
)
|
|
|
|
self.assertFalse(result.success)
|
|
self.assertEqual(result.message, "ordinary failure after accepted step")
|
|
self.assertEqual(len(attempts), 2)
|
|
self.assertEqual(caps_after_accepted_retry, [1.0])
|
|
self.assertEqual(result.solver_segments[0].accepted_step_count, 1)
|
|
self.assertEqual(result.solver_segments[0].recoverable_retry_count, 1)
|
|
|
|
def test_transition_restart_does_not_inherit_retry_first_step(self) -> None:
|
|
import numpy as np
|
|
import scipy.integrate
|
|
|
|
attempts: list[dict[str, object]] = []
|
|
transition_pending = True
|
|
|
|
class TransitionAfterRetryBdf:
|
|
def __init__(self, _fun, t0, y0, t_bound, **kwargs):
|
|
self.attempt_index = len(attempts)
|
|
attempts.append(dict(kwargs))
|
|
self.t = float(t0)
|
|
self.y = np.asarray(y0, dtype=float)
|
|
self.t_bound = float(t_bound)
|
|
self.h_abs = min(0.2, float(kwargs["max_step"]))
|
|
self.status = "running"
|
|
|
|
def step(self):
|
|
if self.attempt_index == 0:
|
|
raise RecoverableTrialStateError("retry before transition")
|
|
self.t = self.t_bound
|
|
self.y = np.asarray([self.t], dtype=float)
|
|
self.status = "finished"
|
|
return None
|
|
|
|
def dense_output(self):
|
|
return lambda time: np.asarray([float(time)], dtype=float)
|
|
|
|
def transition_handler(
|
|
previous_time,
|
|
_previous_state,
|
|
current_time,
|
|
_current_state,
|
|
_dense_state,
|
|
):
|
|
nonlocal transition_pending
|
|
if transition_pending and previous_time <= 0.25 <= current_time:
|
|
transition_pending = False
|
|
return StateTransition(time=0.25, state=[0.25])
|
|
return None
|
|
|
|
with patch.object(scipy.integrate, "BDF", TransitionAfterRetryBdf):
|
|
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,
|
|
),
|
|
state_transition_handler=transition_handler,
|
|
)
|
|
|
|
self.assertTrue(result.success, result.message)
|
|
self.assertEqual(len(attempts), 3)
|
|
self.assertNotIn("first_step", attempts[0])
|
|
self.assertEqual(attempts[1]["first_step"], 0.1)
|
|
self.assertEqual(attempts[1]["max_step"], 0.1)
|
|
self.assertNotIn("first_step", attempts[2])
|
|
self.assertEqual(attempts[2]["max_step"], 1.0)
|
|
|
|
def test_recoverable_retry_stops_at_minimum_step(self) -> None:
|
|
import numpy as np
|
|
import scipy.integrate
|
|
|
|
minimum_step = 64.0 * math.ulp(1.0)
|
|
attempts = 0
|
|
|
|
class MinimumStepBdf:
|
|
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.h_abs = 2.0 * minimum_step
|
|
self.status = "running"
|
|
|
|
def step(self):
|
|
raise RecoverableTrialStateError("minimum step reached")
|
|
|
|
with patch.object(scipy.integrate, "BDF", MinimumStepBdf):
|
|
result = integrate_ode(
|
|
rhs=lambda _time, _state: [0.0],
|
|
initial_state=[1.0],
|
|
config=SolveIVPConfig(t_stop=1.0, method="BDF", max_step=1.0),
|
|
cancel_check=lambda: False,
|
|
)
|
|
|
|
self.assertFalse(result.success)
|
|
self.assertEqual(result.message, "minimum step reached")
|
|
self.assertEqual(attempts, 1)
|
|
retry = result.solver_segments[0].recoverable_retries[0]
|
|
self.assertEqual(retry.attempted_step, 2.0 * minimum_step)
|
|
self.assertIsNone(retry.next_max_step)
|
|
|
|
def test_recoverable_retry_has_sixteen_retry_limit(self) -> None:
|
|
import scipy.integrate
|
|
|
|
attempts: list[dict[str, object]] = []
|
|
|
|
class AlwaysFailingConstructorBdf:
|
|
def __init__(self, _fun, _t0, _y0, _t_bound, **kwargs):
|
|
attempts.append(dict(kwargs))
|
|
raise RecoverableTrialStateError("persistent trial failure")
|
|
|
|
with patch.object(
|
|
scipy.integrate,
|
|
"BDF",
|
|
AlwaysFailingConstructorBdf,
|
|
):
|
|
result = integrate_ode(
|
|
rhs=lambda _time, _state: [0.0],
|
|
initial_state=[1.0],
|
|
config=SolveIVPConfig(t_stop=1.0, method="BDF", max_step=1.0),
|
|
cancel_check=lambda: False,
|
|
)
|
|
|
|
self.assertFalse(result.success)
|
|
self.assertEqual(len(attempts), 17)
|
|
diagnostics = result.solver_segments[0].recoverable_retries
|
|
self.assertEqual(len(diagnostics), 17)
|
|
self.assertEqual(
|
|
sum(retry.next_max_step is not None for retry in diagnostics),
|
|
16,
|
|
)
|
|
self.assertIsNone(diagnostics[-1].next_max_step)
|
|
|
|
def test_recoverable_retry_does_not_swallow_cancellation(self) -> None:
|
|
import numpy as np
|
|
import scipy.integrate
|
|
|
|
attempts = 0
|
|
|
|
class CancellingBdf:
|
|
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 IntegrationCancelled
|
|
|
|
with patch.object(scipy.integrate, "BDF", CancellingBdf):
|
|
result = integrate_ode(
|
|
rhs=lambda _time, _state: [0.0],
|
|
initial_state=[1.0],
|
|
config=SolveIVPConfig(t_stop=1.0, method="BDF", max_step=1.0),
|
|
cancel_check=lambda: False,
|
|
)
|
|
|
|
self.assertFalse(result.success)
|
|
self.assertEqual(result.status, "cancelled")
|
|
self.assertEqual(attempts, 1)
|
|
self.assertEqual(result.solver_segments[0].recoverable_retry_count, 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_stepwise_implicit_solver_restarts_callable_jacobian(self) -> None:
|
|
import numpy as np
|
|
import scipy.integrate
|
|
|
|
for method in ("BDF", "Radau"):
|
|
with self.subTest(method=method):
|
|
received: list[dict[str, object]] = []
|
|
|
|
class RecordingJacobian:
|
|
def __init__(self) -> None:
|
|
self.segment_count = 0
|
|
self.observation_count = 0
|
|
|
|
def start_segment(self) -> None:
|
|
self.segment_count += 1
|
|
|
|
def observe(self, _time, _state, _derivative) -> None:
|
|
self.observation_count += 1
|
|
|
|
def __call__(self, _time, _state):
|
|
return [[0.0]]
|
|
|
|
class RecordingSolver:
|
|
def __init__(self, fun, t0, y0, t_bound, **kwargs):
|
|
received.append(kwargs)
|
|
self.fun = fun
|
|
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.fun(self.t, self.y)
|
|
self.t = self.t_bound
|
|
self.status = "finished"
|
|
return None
|
|
|
|
def dense_output(self):
|
|
state = self.y.copy()
|
|
return lambda _time: state.copy()
|
|
|
|
jacobian = RecordingJacobian()
|
|
with patch.object(scipy.integrate, method, RecordingSolver):
|
|
result = integrate_ode(
|
|
rhs=lambda _time, _state: [0.0],
|
|
initial_state=[1.0],
|
|
config=SolveIVPConfig(
|
|
t_start=0.0,
|
|
t_stop=1.0,
|
|
method=method,
|
|
max_step=1.0,
|
|
),
|
|
t_eval=[0.0, 0.4, 1.0],
|
|
breakpoints=[0.4],
|
|
jac_sparsity=[[True]],
|
|
jac=jacobian,
|
|
)
|
|
|
|
self.assertTrue(result.success, result.message)
|
|
self.assertEqual(jacobian.segment_count, 2)
|
|
self.assertEqual(jacobian.observation_count, 2)
|
|
self.assertEqual(len(received), 2)
|
|
self.assertTrue(
|
|
all(options.get("jac") is jacobian for options in received)
|
|
)
|
|
self.assertTrue(
|
|
all("jac_sparsity" not in options for options in received)
|
|
)
|
|
|
|
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)
|
|
transition_segments = [
|
|
segment
|
|
for segment in result.solver_segments
|
|
if segment.state_transition_count
|
|
]
|
|
self.assertEqual(len(transition_segments), 1)
|
|
self.assertEqual(
|
|
transition_segments[0].state_transition_times,
|
|
(event_time,),
|
|
)
|
|
self.assertEqual(
|
|
transition_segments[0].as_dict()["stateTransitionTimes"],
|
|
[event_time],
|
|
)
|
|
|
|
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()
|