优化求解器重试并校正AMESim机械端口
This commit is contained in:
1 parent
971e8f2336
commit
18d9802f03
15 files changed
+383
-166
No files matched your search
@@ -4,6 +4,7 @@ import types
|
||||
import unittest
|
||||
from unittest.mock import patch
|
||||
|
||||
from app.simulation.core.errors import RecoverableTrialStateError
|
||||
from app.simulation.solvers.solver import (
|
||||
SolveIVPConfig,
|
||||
StateTransition,
|
||||
@@ -71,6 +72,81 @@ class IntegrateOdeTests(unittest.TestCase):
|
||||
self.assertEqual(calls[0]["max_step"], 1.0e-3)
|
||||
self.assertNotIn("first_step", calls[0])
|
||||
|
||||
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],
|
||||
cancel_check=lambda: False,
|
||||
)
|
||||
|
||||
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_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,
|
||||
|
||||
Reference in new issue
Block a user