完善AMESim组件界面与仿真求解稳定性

This commit is contained in:
ljz committed 2026-08-02 00:57:48 +08:00
1 parent e7177ab03e
commit 410ef535e8
34 files changed
+3340 -251

No files matched your search

+119
View File
@@ -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()