Files
SystemSimulationApp/tests/test_core_solver.py
T
ljz 971e8f2336 初版:实现 AMESim 机械因果化与事件求解
初步支持 MECMAS21 刚性质量状态归并、端止事件、恢复系数,以及 LSTP 接触和压力流量显式因果化。

已知问题:显式传播仍会重复扫描全网方程,长时刚性仿真性能待优化;自适应积分器遇到越出物理域的试探状态时,尚未实现恢复并缩步重试。
2026-08-03 15:45:48 +08:00

447 lines
15 KiB
Python

import math
import sys
import types
import unittest
from unittest.mock import patch
from app.simulation.solvers.solver import (
SolveIVPConfig,
StateTransition,
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)))
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()