完善AMESim组件界面与仿真求解稳定性
This commit is contained in:
1 parent
e7177ab03e
commit
410ef535e8
34 files changed
+3340
-251
No files matched your search
@@ -1,7 +1,8 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
from typing import Callable, Literal
|
||||
from typing import Callable, Literal, Sequence
|
||||
|
||||
|
||||
CancellationCheck = Callable[[], bool]
|
||||
@@ -51,6 +52,33 @@ def _append_solution_sample(
|
||||
states[index].append(float(value))
|
||||
|
||||
|
||||
def _normalize_breakpoints(
|
||||
config: SolveIVPConfig,
|
||||
breakpoints: Sequence[float] | None,
|
||||
) -> list[float]:
|
||||
"""Return sorted, unique breakpoints strictly inside the integration span."""
|
||||
|
||||
if breakpoints is None or len(breakpoints) == 0:
|
||||
return []
|
||||
if config.t_stop < config.t_start:
|
||||
raise ValueError("Segmented integration requires t_stop to follow t_start.")
|
||||
|
||||
normalized: list[float] = []
|
||||
for raw_breakpoint in breakpoints:
|
||||
breakpoint = float(raw_breakpoint)
|
||||
if not math.isfinite(breakpoint):
|
||||
raise ValueError("Integration breakpoints must be finite numbers.")
|
||||
if config.t_start < breakpoint < config.t_stop:
|
||||
normalized.append(breakpoint)
|
||||
|
||||
normalized.sort()
|
||||
return [
|
||||
breakpoint
|
||||
for index, breakpoint in enumerate(normalized)
|
||||
if index == 0 or breakpoint != normalized[index - 1]
|
||||
]
|
||||
|
||||
|
||||
def _runge_kutta_4(
|
||||
rhs: Callable[[float, list[float]], list[float]],
|
||||
initial_state: list[float],
|
||||
@@ -114,6 +142,137 @@ def _runge_kutta_4(
|
||||
)
|
||||
|
||||
|
||||
def _runge_kutta_4_segmented(
|
||||
rhs: Callable[[float, list[float]], list[float]],
|
||||
initial_state: list[float],
|
||||
config: SolveIVPConfig,
|
||||
t_eval: list[float] | None,
|
||||
breakpoints: Sequence[float],
|
||||
cancel_check: CancellationCheck | None = None,
|
||||
accepted_step_callback: AcceptedStepCallback | None = None,
|
||||
) -> ODESolution:
|
||||
"""RK4 fallback that never evaluates a pre-breakpoint step at the breakpoint."""
|
||||
|
||||
if t_eval is None:
|
||||
point_count = max(
|
||||
2,
|
||||
int((config.t_stop - config.t_start) / max(config.max_step, 1e-6)) + 1,
|
||||
)
|
||||
sample_step = (config.t_stop - config.t_start) / (point_count - 1)
|
||||
sample_times = [
|
||||
config.t_start + index * sample_step for index in range(point_count)
|
||||
]
|
||||
else:
|
||||
sample_times = [float(time) for time in t_eval]
|
||||
|
||||
state = [float(value) for value in initial_state]
|
||||
states = [[value] for value in state]
|
||||
times = [float(config.t_start)]
|
||||
current_time = float(config.t_start)
|
||||
sample_index = 0
|
||||
while (
|
||||
sample_index < len(sample_times)
|
||||
and sample_times[sample_index] <= config.t_start + 1e-12
|
||||
):
|
||||
sample_index += 1
|
||||
|
||||
status: IntegrationStatus = "completed"
|
||||
message = "Integrated with built-in RK4 fallback because SciPy is unavailable."
|
||||
error: Exception | None = None
|
||||
last_reported_step: float | None = None
|
||||
|
||||
def report_step(time: float) -> None:
|
||||
nonlocal last_reported_step
|
||||
if accepted_step_callback is None:
|
||||
return
|
||||
if last_reported_step is not None and time <= last_reported_step:
|
||||
return
|
||||
accepted_step_callback(float(time))
|
||||
last_reported_step = float(time)
|
||||
|
||||
def advance_to(
|
||||
target_time: float, reported_terminal_time: float | None = None
|
||||
) -> None:
|
||||
nonlocal current_time, state
|
||||
while current_time < target_time - 1e-15:
|
||||
if cancel_check is not None and cancel_check():
|
||||
raise _IntegrationCancelled
|
||||
dt = min(config.max_step, target_time - current_time)
|
||||
k1 = rhs(current_time, state)
|
||||
k2 = rhs(
|
||||
current_time + 0.5 * dt,
|
||||
_vector_add(state, k1, 0.5 * dt),
|
||||
)
|
||||
k3 = rhs(
|
||||
current_time + 0.5 * dt,
|
||||
_vector_add(state, k2, 0.5 * dt),
|
||||
)
|
||||
k4 = rhs(current_time + dt, _vector_add(state, k3, dt))
|
||||
state = [
|
||||
value + (dt / 6.0) * (a + 2.0 * b + 2.0 * c + d)
|
||||
for value, a, b, c, d in zip(state, k1, k2, k3, k4)
|
||||
]
|
||||
current_time += dt
|
||||
report_time = current_time
|
||||
if (
|
||||
reported_terminal_time is not None
|
||||
and current_time >= target_time - 1e-15
|
||||
):
|
||||
report_time = reported_terminal_time
|
||||
report_step(report_time)
|
||||
|
||||
try:
|
||||
segment_ends = [*breakpoints, float(config.t_stop)]
|
||||
for segment_index, segment_end in enumerate(segment_ends):
|
||||
is_breakpoint = segment_index < len(breakpoints)
|
||||
integration_end = (
|
||||
math.nextafter(segment_end, -math.inf) if is_breakpoint else segment_end
|
||||
)
|
||||
|
||||
while (
|
||||
sample_index < len(sample_times)
|
||||
and sample_times[sample_index] <= integration_end
|
||||
):
|
||||
sample_time = float(sample_times[sample_index])
|
||||
advance_to(sample_time)
|
||||
_append_solution_sample(times, states, sample_time, state)
|
||||
sample_index += 1
|
||||
|
||||
advance_to(
|
||||
integration_end,
|
||||
segment_end if is_breakpoint else None,
|
||||
)
|
||||
|
||||
if is_breakpoint:
|
||||
current_time = float(segment_end)
|
||||
report_step(current_time)
|
||||
while (
|
||||
sample_index < len(sample_times)
|
||||
and sample_times[sample_index] <= segment_end
|
||||
):
|
||||
sample_time = float(sample_times[sample_index])
|
||||
_append_solution_sample(times, states, sample_time, state)
|
||||
sample_index += 1
|
||||
except _IntegrationCancelled:
|
||||
status = "cancelled"
|
||||
message = "Simulation was stopped before reaching the requested end time."
|
||||
_append_solution_sample(times, states, current_time, state)
|
||||
except Exception as exc:
|
||||
status = "failed"
|
||||
message = str(exc)
|
||||
error = exc
|
||||
_append_solution_sample(times, states, current_time, state)
|
||||
|
||||
return ODESolution(
|
||||
t=times,
|
||||
y=states,
|
||||
success=status == "completed",
|
||||
message=message,
|
||||
status=status,
|
||||
error=error,
|
||||
)
|
||||
|
||||
|
||||
def _integrate_scipy_stepwise(
|
||||
rhs: Callable[[float, list[float]], list[float]],
|
||||
initial_state: list[float],
|
||||
@@ -121,6 +280,7 @@ def _integrate_scipy_stepwise(
|
||||
t_eval: list[float] | None,
|
||||
cancel_check: CancellationCheck,
|
||||
accepted_step_callback: AcceptedStepCallback | None,
|
||||
breakpoints: Sequence[float] = (),
|
||||
) -> ODESolution:
|
||||
import numpy as np
|
||||
from scipy.integrate import BDF, DOP853, LSODA, RK23, RK45, Radau
|
||||
@@ -141,7 +301,7 @@ def _integrate_scipy_stepwise(
|
||||
states = [[float(value)] for value in initial_state]
|
||||
last_accepted_time = float(config.t_start)
|
||||
last_accepted_state = [float(value) for value in initial_state]
|
||||
sample_times = list(t_eval or [])
|
||||
sample_times = [float(time) for time in (t_eval or [])]
|
||||
sample_index = 0
|
||||
while (
|
||||
sample_index < len(sample_times)
|
||||
@@ -154,96 +314,160 @@ def _integrate_scipy_stepwise(
|
||||
raise _IntegrationCancelled
|
||||
return rhs(float(time), [float(value) for value in state])
|
||||
|
||||
if cancel_check():
|
||||
return ODESolution(
|
||||
t=times,
|
||||
y=states,
|
||||
success=False,
|
||||
message="Simulation was stopped before integration started.",
|
||||
status="cancelled",
|
||||
)
|
||||
|
||||
solver_options = {
|
||||
"rtol": config.rtol,
|
||||
"atol": config.atol,
|
||||
"max_step": config.max_step,
|
||||
}
|
||||
if config.first_step is not None:
|
||||
solver_options["first_step"] = config.first_step
|
||||
|
||||
try:
|
||||
solver = solver_type(
|
||||
cancellable_rhs,
|
||||
config.t_start,
|
||||
np.asarray(initial_state, dtype=float),
|
||||
config.t_stop,
|
||||
**solver_options,
|
||||
)
|
||||
except _IntegrationCancelled:
|
||||
return ODESolution(
|
||||
t=times,
|
||||
y=states,
|
||||
success=False,
|
||||
message="Simulation was stopped before integration started.",
|
||||
status="cancelled",
|
||||
)
|
||||
except Exception as exc:
|
||||
return ODESolution(
|
||||
t=times,
|
||||
y=states,
|
||||
success=False,
|
||||
message=str(exc),
|
||||
status="failed",
|
||||
error=exc,
|
||||
)
|
||||
|
||||
status: IntegrationStatus = "completed"
|
||||
message = "The solver successfully reached the end of the integration interval."
|
||||
error: Exception | None = None
|
||||
last_reported_step: float | None = None
|
||||
|
||||
while solver.status == "running":
|
||||
def report_step(time: float) -> None:
|
||||
nonlocal last_reported_step
|
||||
if accepted_step_callback is None:
|
||||
return
|
||||
if last_reported_step is not None and time <= last_reported_step:
|
||||
return
|
||||
accepted_step_callback(float(time))
|
||||
last_reported_step = float(time)
|
||||
|
||||
segment_ends = [*breakpoints, float(config.t_stop)]
|
||||
for segment_index, segment_end in enumerate(segment_ends):
|
||||
if cancel_check():
|
||||
status = "cancelled"
|
||||
message = "Simulation was stopped before reaching the requested end time."
|
||||
break
|
||||
try:
|
||||
step_message = solver.step()
|
||||
except _IntegrationCancelled:
|
||||
status = "cancelled"
|
||||
message = "Simulation was stopped before reaching the requested end time."
|
||||
break
|
||||
except Exception as exc:
|
||||
status = "failed"
|
||||
message = str(exc)
|
||||
error = exc
|
||||
break
|
||||
|
||||
if solver.status == "failed":
|
||||
status = "failed"
|
||||
message = str(step_message or "Integration step failed.")
|
||||
break
|
||||
|
||||
last_accepted_time = float(solver.t)
|
||||
last_accepted_state = [float(value) for value in solver.y]
|
||||
if sample_times:
|
||||
dense_output = solver.dense_output()
|
||||
while (
|
||||
sample_index < len(sample_times)
|
||||
and sample_times[sample_index] <= last_accepted_time + 1e-12
|
||||
):
|
||||
sample_time = float(sample_times[sample_index])
|
||||
sample_state = [float(value) for value in dense_output(sample_time)]
|
||||
_append_solution_sample(times, states, sample_time, sample_state)
|
||||
sample_index += 1
|
||||
else:
|
||||
_append_solution_sample(
|
||||
times,
|
||||
states,
|
||||
last_accepted_time,
|
||||
last_accepted_state,
|
||||
message = (
|
||||
"Simulation was stopped before integration started."
|
||||
if segment_index == 0
|
||||
else "Simulation was stopped before reaching the requested end time."
|
||||
)
|
||||
if accepted_step_callback is not None:
|
||||
accepted_step_callback(last_accepted_time)
|
||||
break
|
||||
|
||||
is_breakpoint = segment_index < len(breakpoints)
|
||||
integration_end = (
|
||||
math.nextafter(segment_end, -math.inf) if is_breakpoint else segment_end
|
||||
)
|
||||
has_integration_interval = integration_end > last_accepted_time
|
||||
|
||||
if has_integration_interval:
|
||||
solver_options = {
|
||||
"rtol": config.rtol,
|
||||
"atol": config.atol,
|
||||
"max_step": config.max_step,
|
||||
}
|
||||
if config.first_step is not None:
|
||||
solver_options["first_step"] = min(
|
||||
config.first_step,
|
||||
integration_end - last_accepted_time,
|
||||
)
|
||||
|
||||
try:
|
||||
solver = solver_type(
|
||||
cancellable_rhs,
|
||||
last_accepted_time,
|
||||
np.asarray(last_accepted_state, dtype=float),
|
||||
integration_end,
|
||||
**solver_options,
|
||||
)
|
||||
except _IntegrationCancelled:
|
||||
status = "cancelled"
|
||||
message = (
|
||||
"Simulation was stopped before integration started."
|
||||
if segment_index == 0
|
||||
else "Simulation was stopped before reaching the requested end time."
|
||||
)
|
||||
break
|
||||
except Exception as exc:
|
||||
status = "failed"
|
||||
message = str(exc)
|
||||
error = exc
|
||||
break
|
||||
|
||||
while solver.status == "running":
|
||||
if cancel_check():
|
||||
status = "cancelled"
|
||||
message = (
|
||||
"Simulation was stopped before reaching the requested end time."
|
||||
)
|
||||
break
|
||||
try:
|
||||
step_message = solver.step()
|
||||
except _IntegrationCancelled:
|
||||
status = "cancelled"
|
||||
message = (
|
||||
"Simulation was stopped before reaching the requested end time."
|
||||
)
|
||||
break
|
||||
except Exception as exc:
|
||||
status = "failed"
|
||||
message = str(exc)
|
||||
error = exc
|
||||
break
|
||||
|
||||
if solver.status == "failed":
|
||||
status = "failed"
|
||||
message = str(step_message or "Integration step failed.")
|
||||
break
|
||||
|
||||
last_accepted_time = float(solver.t)
|
||||
last_accepted_state = [float(value) for value in solver.y]
|
||||
reported_time = (
|
||||
float(segment_end)
|
||||
if is_breakpoint and solver.status == "finished"
|
||||
else last_accepted_time
|
||||
)
|
||||
if sample_times:
|
||||
dense_output = solver.dense_output()
|
||||
while (
|
||||
sample_index < len(sample_times)
|
||||
and sample_times[sample_index] <= last_accepted_time
|
||||
):
|
||||
sample_time = float(sample_times[sample_index])
|
||||
sample_state = [
|
||||
float(value) for value in dense_output(sample_time)
|
||||
]
|
||||
_append_solution_sample(
|
||||
times,
|
||||
states,
|
||||
sample_time,
|
||||
sample_state,
|
||||
)
|
||||
sample_index += 1
|
||||
else:
|
||||
_append_solution_sample(
|
||||
times,
|
||||
states,
|
||||
reported_time,
|
||||
last_accepted_state,
|
||||
)
|
||||
report_step(reported_time)
|
||||
|
||||
if status != "completed":
|
||||
break
|
||||
|
||||
if is_breakpoint:
|
||||
# The old equation is integrated only to the representable point just
|
||||
# left of the event. The continuous state is then lifted to the exact
|
||||
# event time, where the freshly constructed next solver sees the new
|
||||
# equation immediately.
|
||||
last_accepted_time = float(segment_end)
|
||||
if sample_times:
|
||||
while (
|
||||
sample_index < len(sample_times)
|
||||
and sample_times[sample_index] <= segment_end
|
||||
):
|
||||
sample_time = float(sample_times[sample_index])
|
||||
_append_solution_sample(
|
||||
times,
|
||||
states,
|
||||
sample_time,
|
||||
last_accepted_state,
|
||||
)
|
||||
sample_index += 1
|
||||
elif not has_integration_interval:
|
||||
_append_solution_sample(
|
||||
times,
|
||||
states,
|
||||
last_accepted_time,
|
||||
last_accepted_state,
|
||||
)
|
||||
report_step(last_accepted_time)
|
||||
|
||||
if status != "completed":
|
||||
_append_solution_sample(
|
||||
@@ -270,8 +494,14 @@ def integrate_ode(
|
||||
t_eval: list[float] | None = None,
|
||||
cancel_check: CancellationCheck | None = None,
|
||||
accepted_step_callback: AcceptedStepCallback | None = None,
|
||||
breakpoints: Sequence[float] | None = None,
|
||||
):
|
||||
"""Thin wrapper around scipy.integrate.solve_ivp with a pure-Python fallback."""
|
||||
"""Integrate an ODE, optionally restarting at equation discontinuities.
|
||||
|
||||
Breakpoints are interpreted as right-continuous equation changes: the old
|
||||
equation is integrated to the floating-point left limit, then a fresh solver
|
||||
starts at the exact breakpoint with the unchanged continuous state.
|
||||
"""
|
||||
|
||||
if abs(config.t_stop - config.t_start) <= 1e-15:
|
||||
return ODESolution(
|
||||
@@ -281,9 +511,21 @@ def integrate_ode(
|
||||
message="Skipped integration because t_start equals t_stop.",
|
||||
)
|
||||
|
||||
normalized_breakpoints = _normalize_breakpoints(config, breakpoints)
|
||||
|
||||
try:
|
||||
from scipy.integrate import solve_ivp
|
||||
except ImportError:
|
||||
if normalized_breakpoints:
|
||||
return _runge_kutta_4_segmented(
|
||||
rhs,
|
||||
initial_state,
|
||||
config,
|
||||
t_eval,
|
||||
normalized_breakpoints,
|
||||
cancel_check,
|
||||
accepted_step_callback,
|
||||
)
|
||||
return _runge_kutta_4(
|
||||
rhs,
|
||||
initial_state,
|
||||
@@ -293,14 +535,15 @@ def integrate_ode(
|
||||
accepted_step_callback,
|
||||
)
|
||||
|
||||
if cancel_check is not None:
|
||||
if cancel_check is not None or normalized_breakpoints:
|
||||
return _integrate_scipy_stepwise(
|
||||
rhs,
|
||||
initial_state,
|
||||
config,
|
||||
t_eval,
|
||||
cancel_check,
|
||||
cancel_check or (lambda: False),
|
||||
accepted_step_callback,
|
||||
normalized_breakpoints,
|
||||
)
|
||||
|
||||
solve_options = {
|
||||
|
||||
Reference in new issue
Block a user