完善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

+331 -88
View File
@@ -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 = {