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

@@ -26,7 +26,7 @@ class AmesimPnpl01(AlgebraicComponent):
label="PNPL01 零气动流边界",
library_id="amesim",
category_id="boundary",
symbol="generic",
symbol="amesim_pnpl01",
ports=(PortDisplaySpec("port_1", "left", order=10),),
order=10,
)
@@ -105,7 +105,7 @@ class AmesimPnor001(AlgebraicComponent):
label="PNOR001 常系数气动孔口",
library_id="amesim",
category_id="flow",
symbol="orifice",
symbol="amesim_pnor001",
ports=(
PortDisplaySpec("port_1", "left", order=10),
PortDisplaySpec("port_2", "right", order=20),
@@ -412,7 +412,7 @@ class AmesimPnvo001FixedOpening(AlgebraicComponent):
label="PNVO001 固定开度气动孔口",
library_id="amesim",
category_id="flow",
symbol="orifice",
symbol="amesim_pnvo001_fixed",
ports=(
PortDisplaySpec("port_2", "left", order=10),
PortDisplaySpec("port_3", "right", order=20),
@@ -633,7 +633,7 @@ class AmesimPnvo001SignalOpening(AmesimPnvo001FixedOpening):
label="PNVO001 信号开度气动孔口",
library_id="amesim",
category_id="flow",
symbol="orifice",
symbol="amesim_pnvo001",
ports=(
PortDisplaySpec("res", "left", order=5),
PortDisplaySpec("port_2", "left", order=10),
@@ -105,7 +105,7 @@ class AmesimPnl00r(AlgebraicComponent):
label="PNL00R 气动管路阻力",
library_id="amesim",
category_id="flow",
symbol="pipe",
symbol="amesim_pnl00r",
ports=(
PortDisplaySpec("port_1", "left", order=10),
PortDisplaySpec("port_2", "right", order=20),
@@ -440,7 +440,7 @@ class AmesimPnl0001(ThermodynamicVolumeComponent):
label="PNL0001 C-R 动态管路",
library_id="amesim",
category_id="flow",
symbol="pipe",
symbol="amesim_pnl0001",
ports=(
PortDisplaySpec("port_1", "left", order=10),
PortDisplaySpec("port_2", "right", order=20),
@@ -720,7 +720,7 @@ class AmesimPnl0002(AmesimPnl0001):
label="PNL0002 R-C-R 动态管路",
library_id="amesim",
category_id="flow",
symbol="pipe",
symbol="amesim_pnl0002",
ports=(
PortDisplaySpec("port_1", "left", order=10),
PortDisplaySpec("port_2", "right", order=20),
@@ -920,7 +920,7 @@ class AmesimPnl0003(DynamicComponent):
label="PNL0003 C-R-C 动态管路",
library_id="amesim",
category_id="flow",
symbol="pipe",
symbol="amesim_pnl0003",
ports=(
PortDisplaySpec("port_1", "left", order=10),
PortDisplaySpec("port_2", "right", order=20),
@@ -91,7 +91,7 @@ class AmesimPn3Node2(_AmesimPneumaticNode):
label="PN3NODE2 三端气动节点",
library_id="amesim",
category_id="junctions",
symbol="tee",
symbol="amesim_pn3node2",
ports=(
PortDisplaySpec("port_1", "left", order=10),
PortDisplaySpec("port_2", "right", order=20),
@@ -128,7 +128,7 @@ class AmesimP4Node2(_AmesimPneumaticNode):
label="P4NODE2 四端气动节点",
library_id="amesim",
category_id="junctions",
symbol="generic",
symbol="amesim_p4node2",
ports=(
PortDisplaySpec("port_1", "left", order=10),
PortDisplaySpec("port_2", "right", order=20),
@@ -22,7 +22,7 @@ class AmesimF000(AlgebraicComponent):
label="F000 零力源",
library_id="amesim",
category_id="mechanical",
symbol="generic",
symbol="amesim_f000",
ports=(PortDisplaySpec("port_1", "right", order=10),),
order=10,
)
@@ -73,7 +73,7 @@ class AmesimForc(AlgebraicComponent):
label="FORC 信号转力",
library_id="amesim",
category_id="mechanical",
symbol="signal",
symbol="amesim_forc",
ports=(
PortDisplaySpec("res", "left", order=10),
PortDisplaySpec("port_2", "right", order=20),
@@ -167,7 +167,7 @@ class AmesimMecmas21(DynamicComponent):
label="MECMAS21 一维质量",
library_id="amesim",
category_id="mechanical",
symbol="generic",
symbol="amesim_mecmas21",
ports=(
PortDisplaySpec("port_1", "left", order=10),
PortDisplaySpec("port_2", "right", order=20),
@@ -323,7 +323,7 @@ class AmesimLstp00a(AlgebraicComponent):
label="LSTP00A 弹性接触",
library_id="amesim",
category_id="mechanical",
symbol="generic",
symbol="amesim_lstp00a",
ports=(
PortDisplaySpec("port_1", "left", order=10),
PortDisplaySpec("port_2", "right", order=20),
@@ -425,7 +425,7 @@ class AmesimLmechn1(AlgebraicComponent):
label="LMECHN1 线性机械节点",
library_id="amesim",
category_id="mechanical",
symbol="junction",
symbol="amesim_lmechn1",
ports=tuple(
[PortDisplaySpec(f"port_{index}", "left", order=index * 10) for index in range(1, 9)]
+ [PortDisplaySpec("port_9", "right", order=90)]
@@ -1,6 +1,7 @@
from __future__ import annotations
from collections.abc import Mapping
from math import floor
from app.simulation.core.base import AlgebraicComponent
from app.simulation.core.catalog import ComponentDisplaySpec, PortDisplaySpec
@@ -27,7 +28,7 @@ class AmesimStep0(AlgebraicComponent):
label="STEP0 阶跃信号",
library_id="amesim",
category_id="signals",
symbol="signal",
symbol="amesim_step0",
ports=(PortDisplaySpec("out", "right", order=10),),
order=10,
)
@@ -71,6 +72,15 @@ class AmesimStep0(AlgebraicComponent):
def signal_output_values(self, time: float) -> dict[str, float]:
return {"out": self.output_at(time)}
def signal_event_times(
self,
start_time: float,
stop_time: float,
) -> tuple[float, ...]:
"""Expose the exact STEP0 switch time as an integration split point."""
return (self.time,) if start_time < self.time < stop_time else ()
def component_result_values(self) -> Mapping[str, float]:
return {"y": self.out.signal}
@@ -117,7 +127,7 @@ class AmesimUd00(AlgebraicComponent):
label="UD00 分段线性信号",
library_id="amesim",
category_id="signals",
symbol="signal",
symbol="amesim_ud00",
ports=(PortDisplaySpec("out", "right", order=10),),
order=20,
)
@@ -200,5 +210,58 @@ class AmesimUd00(AlgebraicComponent):
def signal_output_values(self, time: float) -> dict[str, float]:
return {"out": self.output_at(time)}
def signal_event_times(
self,
start_time: float,
stop_time: float,
) -> tuple[float, ...]:
"""Return UD00 start, stage, and repeated cycle boundaries.
The final non-cyclic stage is intentionally not given an end event:
``output_at`` continues that stage's slope after its configured duration.
"""
if stop_time <= start_time:
return ()
active_durations = self.durations[: self.nstages]
stage_offsets = [0.0]
elapsed = 0.0
for duration in active_durations[:-1]:
elapsed += duration
stage_offsets.append(elapsed)
if not self.iscyclic:
return tuple(
sorted(
{
event_time
for offset in stage_offsets
if start_time
< (event_time := self.tstart + offset)
< stop_time
}
)
)
cycle_duration = sum(active_durations)
if cycle_duration <= 0.0:
return ()
events: set[float] = set()
for offset in stage_offsets:
first_boundary = self.tstart + offset
cycle_index = max(
0,
floor((start_time - first_boundary) / cycle_duration) + 1,
)
event_time = first_boundary + cycle_index * cycle_duration
while event_time < stop_time:
if event_time > start_time:
events.add(event_time)
cycle_index += 1
event_time = first_boundary + cycle_index * cycle_duration
return tuple(sorted(events))
def component_result_values(self) -> Mapping[str, float]:
return {"y": self.out.signal}
@@ -100,7 +100,7 @@ class AmesimPnch023(ThermodynamicVolumeComponent):
label="PNCH023 固定容积气室",
library_id="amesim",
category_id="storage",
symbol="tank",
symbol="amesim_pnch023",
ports=(
PortDisplaySpec("port_1", "left", order=10),
PortDisplaySpec("port_2", "right", order=20),
@@ -340,7 +340,7 @@ class AmesimPnch012(ThermodynamicVolumeComponent):
label="PNCH012 变容气室",
library_id="amesim",
category_id="storage",
symbol="tank",
symbol="amesim_pnch012",
ports=(
PortDisplaySpec("port_1", "left", order=10),
PortDisplaySpec("port_2", "right", order=20),
+257 -43
View File
@@ -1,7 +1,7 @@
from __future__ import annotations
from dataclasses import dataclass
from math import sqrt
from math import isfinite, sqrt
from app.simulation.core.ports import PortState, VariableRole
from app.simulation.systems.network import SimulationNetwork
@@ -68,6 +68,7 @@ class PressureFlowSolver:
self.residual_tolerance = residual_tolerance
self.max_evaluations = max_evaluations
self.unknowns = self._build_unknowns()
self._unknowns_by_id = {unknown.id: unknown for unknown in self.unknowns}
self.last_diagnostics: AlgebraicSolveDiagnostics | None = None
def _build_unknowns(self) -> tuple[AlgebraicUnknown, ...]:
@@ -91,48 +92,220 @@ class PressureFlowSolver:
)
return tuple(unknowns)
def _seed_equal_pressures(self) -> None:
for _ in range(max(2, len(self.network.connections))):
changed = False
for connection in self.network.connections:
if connection.kind != "physical":
continue
first = self.network.components[
connection.endpoint_a.component
].get_port(connection.endpoint_a.port)
second = self.network.components[
connection.endpoint_b.component
].get_port(connection.endpoint_b.port)
if first.p > 0.0 and second.p <= 0.0:
second.p = first.p
changed = True
elif second.p > 0.0 and first.p <= 0.0:
first.p = second.p
changed = True
@staticmethod
def _port_key(variable: str, expected_variable: str) -> tuple[str, str] | None:
try:
component_name, port_name, variable_name = variable.rsplit(".", 2)
except ValueError:
return None
if variable_name != expected_variable:
return None
return component_name, port_name
for component in self.network.components.values():
equal_pressure_equations = [
equation
for equation in component.pressure_flow_equation_residuals()
if equation.relation == "equal" and equation.role == "effort"
def _seed_equal_pressures(self) -> None:
"""Lift current state pressures across their complete equality groups.
Dynamic components refresh their own pressure ports before each closure,
while connected algebraic ports retain values from the preceding RHS
evaluation. Merely filling non-positive pressures therefore leaves a
stale, and sometimes badly conditioned, nonlinear initial guess. State
equations expose the current pressure as ``port.p - target``; use that
target as the authoritative anchor for every connected/equal port.
"""
pressure_unknowns = {
(unknown.component, unknown.port): unknown
for unknown in self.unknowns
if unknown.variable == "p"
}
if not pressure_unknowns:
return
parent = {key: key for key in pressure_unknowns}
def find(key: tuple[str, str]) -> tuple[str, str]:
root = key
while parent[root] != root:
root = parent[root]
while parent[key] != key:
next_key = parent[key]
parent[key] = root
key = next_key
return root
def union(first: tuple[str, str], second: tuple[str, str]) -> None:
first_root = find(first)
second_root = find(second)
if first_root != second_root:
parent[second_root] = first_root
for connection in self.network.connections:
if connection.kind != "physical":
continue
first = connection.endpoint_a.key
second = connection.endpoint_b.key
if first in pressure_unknowns and second in pressure_unknowns:
union(first, second)
component_equations = {
component.name: component.pressure_flow_equation_residuals()
for component in self.network.components.values()
}
for equations in component_equations.values():
for equation in equations:
if equation.relation != "equal" or equation.role != "effort":
continue
endpoints = [
endpoint
for variable in equation.variables
if (
(endpoint := self._port_key(variable, "p"))
in pressure_unknowns
)
]
for equation in equal_pressure_equations:
states = []
for variable in equation.variables:
_, port_name, variable_name = variable.rsplit(".", 2)
if variable_name == "p":
states.append(component.get_port(port_name))
if len(states) != 2:
continue
first, second = states
if first.p > 0.0 and second.p <= 0.0:
second.p = first.p
changed = True
elif second.p > 0.0 and first.p <= 0.0:
first.p = second.p
changed = True
if not changed:
break
for endpoint in endpoints[1:]:
union(endpoints[0], endpoint)
members_by_root: dict[tuple[str, str], list[tuple[str, str]]] = {}
for endpoint in pressure_unknowns:
members_by_root.setdefault(find(endpoint), []).append(endpoint)
anchors_by_root: dict[tuple[str, str], list[float]] = {}
for equations in component_equations.values():
for equation in equations:
if equation.relation != "state" or equation.role != "effort":
continue
endpoints = [
endpoint
for variable in equation.variables
if (
(endpoint := self._port_key(variable, "p"))
in pressure_unknowns
)
]
if len(endpoints) != 1:
continue
endpoint = endpoints[0]
unknown = pressure_unknowns[endpoint]
target_pressure = unknown.read() - float(equation.value)
if not isfinite(target_pressure):
continue
# Keep the state-owned port current even when an invalid model
# has conflicting storage anchors in one equality group.
unknown.write(target_pressure)
anchors_by_root.setdefault(find(endpoint), []).append(target_pressure)
for root, members in members_by_root.items():
anchors = anchors_by_root.get(root, [])
if anchors:
pressure_scale = max([abs(value) for value in anchors] + [1.0])
if max(anchors) - min(anchors) > 1.0e-9 * pressure_scale:
# A conflicting multi-storage group is structurally invalid;
# leave it for the residual solver/preparation diagnostics.
continue
target_pressure = sum(anchors) / len(anchors)
for endpoint in members:
pressure_unknowns[endpoint].write(target_pressure)
continue
positive_seed = next(
(
pressure_unknowns[endpoint].read()
for endpoint in members
if pressure_unknowns[endpoint].read() > 0.0
),
None,
)
if positive_seed is None:
continue
for endpoint in members:
unknown = pressure_unknowns[endpoint]
if unknown.read() <= 0.0:
unknown.write(positive_seed)
def _seed_explicit_mass_flows(self) -> None:
"""Initialize explicit ``m_flow - f(...)`` constitutive relations.
AMESim orifices and quasi-steady pneumatic lines expose one mass-flow
unknown with unit coefficient. Once pressure anchors are current, a
residual correction places that flow directly on its constitutive
surface and avoids asking the nonlinear optimizer to discover the
square-root branch from a stale preceding-step value.
"""
seeded_ids: set[str] = set()
for component in self.network.components.values():
for equation in component.pressure_flow_equation_residuals():
if equation.relation != "constitutive" or equation.role != "flow":
continue
mass_flow_unknowns = [
self._unknowns_by_id[variable]
for variable in equation.variables
if variable in self._unknowns_by_id
and self._unknowns_by_id[variable].variable == "m_flow"
]
if len(mass_flow_unknowns) != 1:
continue
unknown = mass_flow_unknowns[0]
target_flow = unknown.read() - float(equation.value)
if not isfinite(target_flow):
continue
unknown.write(target_flow)
seeded_ids.add(unknown.id)
# Complete local two-port balances for explicit elements. Connection
# flow equations remain available to align the adjacent component port.
for component in self.network.components.values():
for equation in component.pressure_flow_equation_residuals():
if equation.relation != "sumToZero" or equation.role != "flow":
continue
mass_flow_unknowns = [
self._unknowns_by_id[variable]
for variable in equation.variables
if variable in self._unknowns_by_id
and self._unknowns_by_id[variable].variable == "m_flow"
]
if len(mass_flow_unknowns) != 2:
continue
seeded = [
unknown for unknown in mass_flow_unknowns if unknown.id in seeded_ids
]
if len(seeded) != 1:
continue
other = next(
unknown for unknown in mass_flow_unknowns if unknown.id not in seeded_ids
)
other.write(-seeded[0].read())
seeded_ids.add(other.id)
# A physical connector imposes the same sum-to-zero flow rule as a
# two-port component. Once an explicit component flow is known, carry
# that guess to the connected storage/boundary port as well. For the
# common volume-orifice-volume topology this makes the seeded state an
# exact algebraic solution and avoids an unnecessary nonlinear solve on
# every ODE/Jacobian evaluation.
for connection in self.network.connections:
if connection.kind != "physical":
continue
endpoint_unknowns = []
for endpoint in connection.endpoints:
unknown = self._unknowns_by_id.get(
f"{endpoint.component}.{endpoint.port}.m_flow"
)
if unknown is not None:
endpoint_unknowns.append(unknown)
if len(endpoint_unknowns) != 2:
continue
seeded = [
unknown for unknown in endpoint_unknowns if unknown.id in seeded_ids
]
if len(seeded) != 1:
continue
other = next(
unknown for unknown in endpoint_unknowns if unknown.id not in seeded_ids
)
other.write(-seeded[0].read())
seeded_ids.add(other.id)
def _scales(self) -> dict[str, float]:
pressure_scale = max(
@@ -184,6 +357,7 @@ class PressureFlowSolver:
) from exc
self._seed_equal_pressures()
self._seed_explicit_mass_flows()
scales = self._scales()
pressure_scale = scales["p"]
flow_scale = scales["m_flow"]
@@ -216,6 +390,40 @@ class PressureFlowSolver:
return pressure_scale
return max([scales.get(name, 1.0) for name in variable_names] + [1.0])
seeded_equations = self.network.pressure_flow_equation_residuals()
seeded_scaled = [
abs(equation.value / equation_scale(equation))
for equation in seeded_equations
]
seeded_max_scaled_residual = max(seeded_scaled, default=0.0)
seeded_unknown_values = [
(unknown, unknown.read()) for unknown in self.unknowns
]
seeded_unknowns_are_feasible = all(
isfinite(value)
and (unknown.variable != "p" or value >= 1.0)
for unknown, value in seeded_unknown_values
)
if (
seeded_unknowns_are_feasible
and all(isfinite(value) for value in seeded_scaled)
and seeded_max_scaled_residual <= self.residual_tolerance
):
diagnostics = AlgebraicSolveDiagnostics(
success=True,
message="Seeded pressure-flow state satisfies the residual tolerance.",
evaluations=0,
pressure_scale=pressure_scale,
flow_scale=flow_scale,
max_scaled_residual=seeded_max_scaled_residual,
max_raw_residual=max(
(abs(item.value) for item in seeded_equations),
default=0.0,
),
)
self.last_diagnostics = diagnostics
return diagnostics
x0 = np.asarray(
[
(
@@ -269,14 +477,20 @@ class PressureFlowSolver:
)
for equation in equations
]
success = bool(result.success) and max(scaled, default=0.0) <= self.residual_tolerance
max_scaled_residual = max(scaled, default=0.0)
residuals_converged = (
all(isfinite(value) for value in scaled)
and max_scaled_residual <= self.residual_tolerance
)
optimizer_status_is_acceptable = bool(result.success) or int(result.status) == 0
success = residuals_converged and optimizer_status_is_acceptable
diagnostics = AlgebraicSolveDiagnostics(
success=success,
message=str(result.message),
evaluations=int(result.nfev),
pressure_scale=pressure_scale,
flow_scale=flow_scale,
max_scaled_residual=max(scaled, default=0.0),
max_scaled_residual=max_scaled_residual,
max_raw_residual=max((abs(item.value) for item in equations), default=0.0),
)
self.last_diagnostics = diagnostics
+47
View File
@@ -1,6 +1,7 @@
from __future__ import annotations
from dataclasses import dataclass
from math import isfinite
from typing import Protocol
from app.simulation.systems.network import Endpoint, SimulationNetwork
@@ -13,6 +14,21 @@ class SignalOutputComponent(Protocol):
...
class SignalEventSource(Protocol):
"""Optional contract for signal sources with known time discontinuities."""
name: str
def signal_event_times(
self,
start_time: float,
stop_time: float,
) -> tuple[float, ...]:
"""Return event times strictly inside ``(start_time, stop_time)``."""
...
@dataclass(frozen=True)
class SignalSolveDiagnostics:
propagated: int
@@ -51,6 +67,37 @@ class SignalResolver:
self.last_diagnostics = diagnostics
return diagnostics
def event_times(self, start_time: float, stop_time: float) -> tuple[float, ...]:
"""Collect optional source events that can be used as integration splits.
Event discovery is deliberately duck typed so existing signal-output
components remain valid without implementing ``signal_event_times``.
"""
start = float(start_time)
stop = float(stop_time)
if not isfinite(start) or not isfinite(stop):
raise ValueError("Signal event interval must be finite.")
if stop < start:
raise ValueError("Signal event interval stop must not precede start.")
if stop == start:
return ()
events: set[float] = set()
for component in self.network.components.values():
source_event_times = getattr(component, "signal_event_times", None)
if source_event_times is None:
continue
for raw_time in source_event_times(start, stop):
event_time = float(raw_time)
if not isfinite(event_time):
raise ValueError(
f"Signal event time from component '{component.name}' must be finite."
)
if start < event_time < stop:
events.add(event_time)
return tuple(sorted(events))
def _source_target(self, endpoints: tuple[Endpoint, Endpoint]) -> tuple[Endpoint, Endpoint]:
first, second = endpoints
first_port = self.network.components[first.component].get_port(first.port)
+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 = {
+6
View File
@@ -327,6 +327,10 @@ class GenericFluidSystem:
report_progress(0.0, "initializing", force=True)
t_eval = simulation_sample_times(config, sample_step)
signal_event_times = self.signal_resolver.event_times(
config.t_start,
config.t_stop,
)
initial_state = self.consistent_initial_state_vector()
report_progress(0.0, "integrating", force=True)
duration = config.t_stop - config.t_start
@@ -356,6 +360,7 @@ class GenericFluidSystem:
accepted_step_callback=(
report_solver_time if cancel_check is not None else None
),
breakpoints=signal_event_times,
)
if isinstance(solution, ODESolution):
run_status: SimulationRunStatus = solution.status
@@ -428,6 +433,7 @@ class GenericFluidSystem:
},
"signal": {
"propagations": self.signal_propagation_count,
"eventTimes": list(signal_event_times),
"last": (
self.signal_resolver.last_diagnostics.as_dict()
if self.signal_resolver.last_diagnostics is not None