完善仿真交互、结果展示与模型元数据
This commit is contained in:
1 parent
f1256a121d
commit
f7f1078911
26 files changed
+9042
-575
No files matched your search
+158
-2
@@ -2,17 +2,29 @@ from __future__ import annotations
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from collections.abc import Mapping
|
||||
from typing import Any
|
||||
from typing import Any, ClassVar
|
||||
|
||||
from PythonModels.core.equations import EquationResidual
|
||||
from PythonModels.core.metadata import (
|
||||
ParameterDefinition,
|
||||
ResultVariableDefinition,
|
||||
ResultVariableMetadata,
|
||||
THERMODYNAMIC_VOLUME_RESULT_VARIABLES,
|
||||
)
|
||||
from PythonModels.core.ports import PortDefinition, PortState
|
||||
|
||||
|
||||
class Component(ABC):
|
||||
MODEL_TYPE: ClassVar[str | None] = None
|
||||
PORTS: ClassVar[tuple[PortDefinition, ...]] = ()
|
||||
PARAMETERS: ClassVar[tuple[ParameterDefinition, ...]] = ()
|
||||
RESULT_VARIABLES: ClassVar[tuple[ResultVariableDefinition, ...]] = ()
|
||||
|
||||
def __init__(self, name: str) -> None:
|
||||
self.name = name
|
||||
self.model_type = self.__class__.__name__.lower()
|
||||
self.model_type = self.MODEL_TYPE or self.__class__.__name__.lower()
|
||||
self._ports: dict[str, PortState] = {}
|
||||
self._parameter_values: dict[str, float] = {}
|
||||
|
||||
@property
|
||||
def ports(self) -> dict[str, PortState]:
|
||||
@@ -35,12 +47,133 @@ class Component(ABC):
|
||||
self._ports[definition.name] = port
|
||||
return port
|
||||
|
||||
def register_declared_port(self, name: str) -> PortState:
|
||||
try:
|
||||
definition = next(item for item in self.PORTS if item.name == name)
|
||||
except StopIteration as exc:
|
||||
raise ValueError(
|
||||
f"Component model {self.model_type} does not declare port {name}."
|
||||
) from exc
|
||||
return self.register_port(PortState(definition=definition))
|
||||
|
||||
def set_parameter_values(self, values: Mapping[str, float]) -> None:
|
||||
definitions = {definition.name: definition for definition in self.PARAMETERS}
|
||||
unknown = sorted(set(values) - set(definitions))
|
||||
if unknown:
|
||||
raise ValueError(
|
||||
f"Component {self.name} contains unsupported parameters: "
|
||||
+ ", ".join(unknown)
|
||||
+ "."
|
||||
)
|
||||
missing = sorted(set(definitions) - set(values))
|
||||
if missing:
|
||||
raise ValueError(
|
||||
f"Component {self.name} is missing parameters: "
|
||||
+ ", ".join(missing)
|
||||
+ "."
|
||||
)
|
||||
|
||||
resolved: dict[str, float] = {}
|
||||
for name, definition in definitions.items():
|
||||
value = float(values[name])
|
||||
message = definition.validation_message(value)
|
||||
if message is not None:
|
||||
raise ValueError(
|
||||
f"Parameter '{name}' on component '{self.name}' {message}."
|
||||
)
|
||||
resolved[name] = value
|
||||
self._parameter_values = resolved
|
||||
|
||||
@property
|
||||
def parameter_values(self) -> dict[str, float]:
|
||||
return dict(self._parameter_values)
|
||||
|
||||
def get_port(self, name: str) -> PortState:
|
||||
try:
|
||||
return self._ports[name]
|
||||
except KeyError as exc:
|
||||
raise ValueError(f"Component {self.name} has no port named {name}.") from exc
|
||||
|
||||
def component_result_values(self) -> Mapping[str, float]:
|
||||
return {}
|
||||
|
||||
def result_values(self) -> dict[str, float]:
|
||||
component_values = dict(self.component_result_values())
|
||||
declared = {definition.name: definition for definition in self.RESULT_VARIABLES}
|
||||
unknown = sorted(set(component_values) - set(declared))
|
||||
if unknown:
|
||||
raise ValueError(
|
||||
f"Component {self.name} returned undeclared result variables: "
|
||||
+ ", ".join(unknown)
|
||||
+ "."
|
||||
)
|
||||
|
||||
values: dict[str, float] = {}
|
||||
for name, definition in declared.items():
|
||||
if not definition.visible:
|
||||
continue
|
||||
if name not in component_values:
|
||||
raise ValueError(
|
||||
f"Component {self.name} did not provide declared result variable {name}."
|
||||
)
|
||||
values[name] = float(component_values[name])
|
||||
|
||||
for port_definition in self.port_definitions:
|
||||
port = self.get_port(port_definition.name)
|
||||
for variable in port_definition.variables:
|
||||
if not variable.result_visible:
|
||||
continue
|
||||
values[f"{port_definition.name}.{variable.name}"] = float(
|
||||
getattr(port, variable.name)
|
||||
)
|
||||
return values
|
||||
|
||||
def result_variable_metadata(self) -> tuple[ResultVariableMetadata, ...]:
|
||||
metadata = [
|
||||
ResultVariableMetadata(
|
||||
key=f"{self.name}.{definition.name}",
|
||||
component_id=self.name,
|
||||
component_type=self.model_type,
|
||||
scope="component",
|
||||
name=definition.name,
|
||||
label=definition.label,
|
||||
quantity=definition.quantity,
|
||||
unit=definition.unit,
|
||||
category=definition.category,
|
||||
order=definition.order,
|
||||
)
|
||||
for definition in self.RESULT_VARIABLES
|
||||
if definition.visible
|
||||
]
|
||||
for port_definition in self.port_definitions:
|
||||
for variable in port_definition.variables:
|
||||
if not variable.result_visible:
|
||||
continue
|
||||
metadata.append(
|
||||
ResultVariableMetadata(
|
||||
key=f"{self.name}.{port_definition.name}.{variable.name}",
|
||||
component_id=self.name,
|
||||
component_type=self.model_type,
|
||||
scope="port",
|
||||
port_name=port_definition.name,
|
||||
name=variable.name,
|
||||
label=variable.label or variable.name,
|
||||
quantity=variable.quantity or variable.name,
|
||||
unit=variable.unit,
|
||||
category=variable.role,
|
||||
order=variable.order,
|
||||
)
|
||||
)
|
||||
return tuple(metadata)
|
||||
|
||||
def parameter_interface_dicts(self) -> list[dict[str, object]]:
|
||||
return [
|
||||
definition.as_interface_dict(
|
||||
value=self._parameter_values.get(definition.name)
|
||||
)
|
||||
for definition in self.PARAMETERS
|
||||
]
|
||||
|
||||
def pressure_flow_equation_residuals(self) -> tuple[EquationResidual, ...]:
|
||||
"""Return algebraic residuals after the network assigns port states."""
|
||||
|
||||
@@ -97,5 +230,28 @@ class DynamicComponent(Component):
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class ThermodynamicVolumeComponent(DynamicComponent):
|
||||
"""Two-state gas volume exposing the shared thermodynamic result contract."""
|
||||
|
||||
RESULT_VARIABLES = THERMODYNAMIC_VOLUME_RESULT_VARIABLES
|
||||
|
||||
def component_result_values(self) -> Mapping[str, float]:
|
||||
state = self.get_state_vector()
|
||||
if len(state) < 2:
|
||||
raise ValueError(
|
||||
f"Thermodynamic component {self.name} must expose mass and energy states."
|
||||
)
|
||||
properties = self.refresh_thermodynamic_ports()
|
||||
return {
|
||||
"m": float(state[0]),
|
||||
"U": float(state[1]),
|
||||
"p": float(properties.p),
|
||||
"T": float(properties.T),
|
||||
"rho": float(properties.rho),
|
||||
"u": float(properties.u),
|
||||
"h": float(properties.h),
|
||||
}
|
||||
|
||||
|
||||
class AlgebraicComponent(Component):
|
||||
"""Stateless element described by algebraic constraints only."""
|
||||
@@ -0,0 +1,156 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from math import isfinite
|
||||
from typing import Literal
|
||||
|
||||
|
||||
ResultVariableScope = Literal["component", "port"]
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ParameterDefinition:
|
||||
"""User-configurable model input expressed in the backend SI contract."""
|
||||
|
||||
name: str
|
||||
default: float
|
||||
label: str = ""
|
||||
quantity: str = "dimensionless"
|
||||
unit: str = ""
|
||||
minimum: float | None = None
|
||||
maximum: float | None = None
|
||||
minimum_exclusive: bool = False
|
||||
|
||||
def validation_message(self, value: float) -> str | None:
|
||||
if not isfinite(value):
|
||||
return "must be finite"
|
||||
if self.minimum is not None:
|
||||
if self.minimum_exclusive and value <= self.minimum:
|
||||
return f"must be greater than {self.minimum:g}"
|
||||
if not self.minimum_exclusive and value < self.minimum:
|
||||
return f"must be at least {self.minimum:g}"
|
||||
if self.maximum is not None and value > self.maximum:
|
||||
return f"must be at most {self.maximum:g}"
|
||||
return None
|
||||
|
||||
def as_interface_dict(self, *, value: float | None = None) -> dict[str, object]:
|
||||
payload: dict[str, object] = {
|
||||
"name": self.name,
|
||||
"label": self.label or self.name,
|
||||
"quantity": self.quantity,
|
||||
"unit": self.unit,
|
||||
"default": self.default,
|
||||
"minimumExclusive": self.minimum_exclusive,
|
||||
}
|
||||
if self.minimum is not None:
|
||||
payload["minimum"] = self.minimum
|
||||
if self.maximum is not None:
|
||||
payload["maximum"] = self.maximum
|
||||
if value is not None:
|
||||
payload["value"] = value
|
||||
return payload
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ResultVariableDefinition:
|
||||
"""Component-relative declaration of a user-visible simulation result."""
|
||||
|
||||
name: str
|
||||
label: str
|
||||
quantity: str
|
||||
unit: str = ""
|
||||
category: str = "derived"
|
||||
order: int = 0
|
||||
visible: bool = True
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ResultVariableMetadata:
|
||||
"""A result declaration bound to one concrete component instance."""
|
||||
|
||||
key: str
|
||||
component_id: str
|
||||
component_type: str
|
||||
scope: ResultVariableScope
|
||||
name: str
|
||||
label: str
|
||||
quantity: str
|
||||
unit: str
|
||||
category: str
|
||||
order: int
|
||||
port_name: str | None = None
|
||||
|
||||
def as_dict(self) -> dict[str, object]:
|
||||
return {
|
||||
"key": self.key,
|
||||
"componentId": self.component_id,
|
||||
"componentType": self.component_type,
|
||||
"scope": self.scope,
|
||||
"portName": self.port_name,
|
||||
"name": self.name,
|
||||
"label": self.label,
|
||||
"quantity": self.quantity,
|
||||
"unit": self.unit,
|
||||
"category": self.category,
|
||||
"order": self.order,
|
||||
}
|
||||
|
||||
|
||||
THERMODYNAMIC_VOLUME_RESULT_VARIABLES = (
|
||||
ResultVariableDefinition(
|
||||
name="m",
|
||||
label="质量",
|
||||
quantity="mass",
|
||||
unit="kg",
|
||||
category="state",
|
||||
order=10,
|
||||
),
|
||||
ResultVariableDefinition(
|
||||
name="U",
|
||||
label="内能",
|
||||
quantity="internal_energy",
|
||||
unit="J",
|
||||
category="state",
|
||||
order=20,
|
||||
),
|
||||
ResultVariableDefinition(
|
||||
name="p",
|
||||
label="压力",
|
||||
quantity="pressure",
|
||||
unit="Pa",
|
||||
category="thermodynamic",
|
||||
order=30,
|
||||
),
|
||||
ResultVariableDefinition(
|
||||
name="T",
|
||||
label="温度",
|
||||
quantity="temperature",
|
||||
unit="K",
|
||||
category="thermodynamic",
|
||||
order=40,
|
||||
),
|
||||
ResultVariableDefinition(
|
||||
name="rho",
|
||||
label="密度",
|
||||
quantity="density",
|
||||
unit="kg/m³",
|
||||
category="thermodynamic",
|
||||
order=50,
|
||||
),
|
||||
ResultVariableDefinition(
|
||||
name="u",
|
||||
label="比内能",
|
||||
quantity="specific_internal_energy",
|
||||
unit="J/kg",
|
||||
category="thermodynamic",
|
||||
order=60,
|
||||
),
|
||||
ResultVariableDefinition(
|
||||
name="h",
|
||||
label="比焓",
|
||||
quantity="specific_enthalpy",
|
||||
unit="J/kg",
|
||||
category="thermodynamic",
|
||||
order=70,
|
||||
),
|
||||
)
|
||||
@@ -4,6 +4,7 @@ from dataclasses import dataclass
|
||||
|
||||
from PythonModels.core.base import Component, DynamicComponent
|
||||
from PythonModels.core.equations import EquationResidual
|
||||
from PythonModels.core.metadata import ResultVariableMetadata
|
||||
from PythonModels.core.ports import PortState
|
||||
|
||||
|
||||
@@ -257,6 +258,13 @@ class SimulationNetwork:
|
||||
if cursor != len(values):
|
||||
raise ValueError("State vector length does not match dynamic components.")
|
||||
|
||||
def result_variable_metadata(self) -> tuple[ResultVariableMetadata, ...]:
|
||||
return tuple(
|
||||
variable
|
||||
for component in self.components.values()
|
||||
for variable in component.result_variable_metadata()
|
||||
)
|
||||
|
||||
def summary(self) -> str:
|
||||
lines = [f"Network: {self.name}", "Components:"]
|
||||
for name, component in self.components.items():
|
||||
@@ -281,10 +289,15 @@ class SimulationNetwork:
|
||||
{
|
||||
"id": component.name,
|
||||
"type": component.model_type,
|
||||
"parameters": component.parameter_interface_dicts(),
|
||||
"ports": [
|
||||
definition.as_interface_dict()
|
||||
for definition in component.port_definitions
|
||||
],
|
||||
"resultVariables": [
|
||||
variable.as_dict()
|
||||
for variable in component.result_variable_metadata()
|
||||
],
|
||||
}
|
||||
for component in self.components.values()
|
||||
],
|
||||
|
||||
@@ -16,12 +16,22 @@ class PortVariableDefinition:
|
||||
name: str
|
||||
role: VariableRole
|
||||
connection_rule: ConnectionRule
|
||||
label: str = field(default="", compare=False)
|
||||
quantity: str = field(default="", compare=False)
|
||||
unit: str = field(default="", compare=False)
|
||||
result_visible: bool = field(default=True, compare=False)
|
||||
order: int = field(default=0, compare=False)
|
||||
|
||||
def as_interface_dict(self) -> dict[str, str]:
|
||||
def as_interface_dict(self) -> dict[str, object]:
|
||||
return {
|
||||
"name": self.name,
|
||||
"role": self.role,
|
||||
"connectionRule": self.connection_rule,
|
||||
"label": self.label or self.name,
|
||||
"quantity": self.quantity or self.name,
|
||||
"unit": self.unit,
|
||||
"resultVisible": self.result_visible,
|
||||
"order": self.order,
|
||||
}
|
||||
|
||||
|
||||
@@ -50,9 +60,33 @@ class PortDefinition:
|
||||
nominal_role=nominal_role,
|
||||
positive_flow_direction="intoComponent",
|
||||
variables=(
|
||||
PortVariableDefinition("p", "effort", "equal"),
|
||||
PortVariableDefinition("m_flow", "flow", "sumToZero"),
|
||||
PortVariableDefinition("h_outflow", "stream", "streamMix"),
|
||||
PortVariableDefinition(
|
||||
"p",
|
||||
"effort",
|
||||
"equal",
|
||||
label="压力",
|
||||
quantity="pressure",
|
||||
unit="Pa",
|
||||
order=10,
|
||||
),
|
||||
PortVariableDefinition(
|
||||
"m_flow",
|
||||
"flow",
|
||||
"sumToZero",
|
||||
label="质量流量",
|
||||
quantity="mass_flow",
|
||||
unit="kg/s",
|
||||
order=20,
|
||||
),
|
||||
PortVariableDefinition(
|
||||
"h_outflow",
|
||||
"stream",
|
||||
"streamMix",
|
||||
label="流出比焓",
|
||||
quantity="specific_enthalpy",
|
||||
unit="J/kg",
|
||||
order=30,
|
||||
),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
+224
-19
@@ -1,7 +1,16 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Callable
|
||||
from typing import Callable, Literal
|
||||
|
||||
|
||||
CancellationCheck = Callable[[], bool]
|
||||
AcceptedStepCallback = Callable[[float], None]
|
||||
IntegrationStatus = Literal["completed", "cancelled", "failed"]
|
||||
|
||||
|
||||
class _IntegrationCancelled(Exception):
|
||||
pass
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
@@ -20,17 +29,34 @@ class ODESolution:
|
||||
y: list[list[float]]
|
||||
success: bool
|
||||
message: str
|
||||
status: IntegrationStatus = "completed"
|
||||
error: Exception | None = None
|
||||
|
||||
|
||||
def _vector_add(a: list[float], b: list[float], scale: float = 1.0) -> list[float]:
|
||||
return [x + scale * y for x, y in zip(a, b)]
|
||||
|
||||
|
||||
def _append_solution_sample(
|
||||
times: list[float],
|
||||
states: list[list[float]],
|
||||
time: float,
|
||||
state: list[float],
|
||||
) -> None:
|
||||
if times and time <= times[-1] + 1e-12:
|
||||
return
|
||||
times.append(float(time))
|
||||
for index, value in enumerate(state):
|
||||
states[index].append(float(value))
|
||||
|
||||
|
||||
def _runge_kutta_4(
|
||||
rhs: Callable[[float, list[float]], list[float]],
|
||||
initial_state: list[float],
|
||||
config: SolveIVPConfig,
|
||||
t_eval: list[float] | None,
|
||||
cancel_check: CancellationCheck | None = None,
|
||||
accepted_step_callback: AcceptedStepCallback | None = None,
|
||||
) -> ODESolution:
|
||||
if t_eval is None:
|
||||
point_count = max(
|
||||
@@ -44,29 +70,189 @@ def _runge_kutta_4(
|
||||
states = [[value] for value in state]
|
||||
times = [float(t_eval[0])]
|
||||
current_time = float(t_eval[0])
|
||||
status: IntegrationStatus = "completed"
|
||||
message = "Integrated with built-in RK4 fallback because SciPy is unavailable."
|
||||
error: Exception | None = None
|
||||
|
||||
for target_time in t_eval[1:]:
|
||||
while current_time < target_time - 1e-15:
|
||||
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
|
||||
try:
|
||||
for target_time in t_eval[1:]:
|
||||
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
|
||||
if accepted_step_callback is not None:
|
||||
accepted_step_callback(current_time)
|
||||
|
||||
times.append(float(target_time))
|
||||
for index, value in enumerate(state):
|
||||
states[index].append(value)
|
||||
_append_solution_sample(times, states, target_time, state)
|
||||
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=True,
|
||||
message="Integrated with built-in RK4 fallback because SciPy is unavailable.",
|
||||
success=status == "completed",
|
||||
message=message,
|
||||
status=status,
|
||||
error=error,
|
||||
)
|
||||
|
||||
|
||||
def _integrate_scipy_stepwise(
|
||||
rhs: Callable[[float, list[float]], list[float]],
|
||||
initial_state: list[float],
|
||||
config: SolveIVPConfig,
|
||||
t_eval: list[float] | None,
|
||||
cancel_check: CancellationCheck,
|
||||
accepted_step_callback: AcceptedStepCallback | None,
|
||||
) -> ODESolution:
|
||||
import numpy as np
|
||||
from scipy.integrate import BDF, DOP853, LSODA, RK23, RK45, Radau
|
||||
|
||||
solver_types = {
|
||||
"BDF": BDF,
|
||||
"DOP853": DOP853,
|
||||
"LSODA": LSODA,
|
||||
"RK23": RK23,
|
||||
"RK45": RK45,
|
||||
"Radau": Radau,
|
||||
}
|
||||
solver_type = solver_types.get(config.method)
|
||||
if solver_type is None:
|
||||
raise ValueError(f"Unsupported integration method: {config.method}")
|
||||
|
||||
times = [float(config.t_start)]
|
||||
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_index = 0
|
||||
while (
|
||||
sample_index < len(sample_times)
|
||||
and sample_times[sample_index] <= config.t_start + 1e-12
|
||||
):
|
||||
sample_index += 1
|
||||
|
||||
def cancellable_rhs(time, state):
|
||||
if cancel_check():
|
||||
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",
|
||||
)
|
||||
|
||||
try:
|
||||
solver = solver_type(
|
||||
cancellable_rhs,
|
||||
config.t_start,
|
||||
np.asarray(initial_state, dtype=float),
|
||||
config.t_stop,
|
||||
rtol=config.rtol,
|
||||
atol=config.atol,
|
||||
max_step=config.max_step,
|
||||
)
|
||||
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
|
||||
|
||||
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]
|
||||
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,
|
||||
)
|
||||
if accepted_step_callback is not None:
|
||||
accepted_step_callback(last_accepted_time)
|
||||
|
||||
if status != "completed":
|
||||
_append_solution_sample(
|
||||
times,
|
||||
states,
|
||||
last_accepted_time,
|
||||
last_accepted_state,
|
||||
)
|
||||
|
||||
return ODESolution(
|
||||
t=times,
|
||||
y=states,
|
||||
success=status == "completed",
|
||||
message=message,
|
||||
status=status,
|
||||
error=error,
|
||||
)
|
||||
|
||||
|
||||
@@ -75,6 +261,8 @@ def integrate_ode(
|
||||
initial_state: list[float],
|
||||
config: SolveIVPConfig,
|
||||
t_eval: list[float] | None = None,
|
||||
cancel_check: CancellationCheck | None = None,
|
||||
accepted_step_callback: AcceptedStepCallback | None = None,
|
||||
):
|
||||
"""Thin wrapper around scipy.integrate.solve_ivp with a pure-Python fallback."""
|
||||
|
||||
@@ -89,7 +277,24 @@ def integrate_ode(
|
||||
try:
|
||||
from scipy.integrate import solve_ivp
|
||||
except ImportError:
|
||||
return _runge_kutta_4(rhs, initial_state, config, t_eval)
|
||||
return _runge_kutta_4(
|
||||
rhs,
|
||||
initial_state,
|
||||
config,
|
||||
t_eval,
|
||||
cancel_check,
|
||||
accepted_step_callback,
|
||||
)
|
||||
|
||||
if cancel_check is not None:
|
||||
return _integrate_scipy_stepwise(
|
||||
rhs,
|
||||
initial_state,
|
||||
config,
|
||||
t_eval,
|
||||
cancel_check,
|
||||
accepted_step_callback,
|
||||
)
|
||||
|
||||
return solve_ivp(
|
||||
fun=rhs,
|
||||
|
||||
Reference in new issue
Block a user