完善AMESim组件界面与仿真求解稳定性
This commit is contained in:
1 parent
e7177ab03e
commit
410ef535e8
34 files changed
+3340
-251
No files matched your search
@@ -1,3 +1,4 @@
|
||||
import math
|
||||
import sys
|
||||
import types
|
||||
import unittest
|
||||
@@ -78,6 +79,124 @@ class IntegrateOdeTests(unittest.TestCase):
|
||||
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()
|
||||
Reference in new issue
Block a user