初版:实现 AMESim 机械因果化与事件求解

初步支持 MECMAS21 刚性质量状态归并、端止事件、恢复系数,以及 LSTP 接触和压力流量显式因果化。

已知问题:显式传播仍会重复扫描全网方程,长时刚性仿真性能待优化;自适应积分器遇到越出物理域的试探状态时,尚未实现恢复并缩步重试。
This commit is contained in:
ljz committed 2026-08-03 15:45:48 +08:00
1 parent de265cdde6
commit 971e8f2336
9 files changed
+2808 -169

No files matched your search

+560 -115
View File
@@ -1,7 +1,7 @@
from __future__ import annotations
from dataclasses import dataclass
from math import isfinite, sqrt
from math import expm1, isfinite, log, sqrt
from app.simulation.core.ports import PortState, VariableRole
from app.simulation.systems.network import SimulationNetwork
@@ -32,6 +32,22 @@ class AlgebraicUnknown:
setattr(self.state, self.variable, float(value))
@dataclass(frozen=True)
class EffortEqualityGroup:
variable: str
members: tuple[AlgebraicUnknown, ...]
anchors: tuple[tuple[AlgebraicUnknown, float], ...]
@dataclass(frozen=True)
class UnilateralContactBinding:
component: object
algebraic_group: EffortEqualityGroup
neighbor_force: AlgebraicUnknown
algebraic_port: int
force_sign: float
@dataclass(frozen=True)
class AlgebraicSolveDiagnostics:
success: bool
@@ -102,26 +118,33 @@ class PressureFlowSolver:
return None
return component_name, port_name
def _seed_equal_pressures(self) -> None:
"""Lift current state pressures across their complete equality groups.
def _seed_equal_efforts(self) -> None:
"""Lift state-owned efforts 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.
Dynamic components refresh their own ports before each closure, while
connected algebraic ports retain values from the preceding RHS
evaluation. State equations expose the current effort as
``port.variable - target``; use that target as the authoritative anchor
for every connected/equal pressure, displacement, and velocity port
before evaluating explicit flow laws.
"""
pressure_unknowns = {
for variable in ("p", "x", "v"):
self._seed_equal_effort(variable)
def _effort_equality_groups(
self,
variable: str,
) -> tuple[EffortEqualityGroup, ...]:
effort_unknowns = {
(unknown.component, unknown.port): unknown
for unknown in self.unknowns
if unknown.variable == "p"
if unknown.variable == variable
}
if not pressure_unknowns:
return
if not effort_unknowns:
return ()
parent = {key: key for key in pressure_unknowns}
parent = {key: key for key in effort_unknowns}
def find(key: tuple[str, str]) -> tuple[str, str]:
root = key
@@ -144,7 +167,7 @@ class PressureFlowSolver:
continue
first = connection.endpoint_a.key
second = connection.endpoint_b.key
if first in pressure_unknowns and second in pressure_unknowns:
if first in effort_unknowns and second in effort_unknowns:
union(first, second)
component_equations = {
@@ -157,155 +180,532 @@ class PressureFlowSolver:
continue
endpoints = [
endpoint
for variable in equation.variables
for equation_variable in equation.variables
if (
(endpoint := self._port_key(variable, "p"))
in pressure_unknowns
(endpoint := self._port_key(equation_variable, variable))
in effort_unknowns
)
]
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)
members_by_root: dict[tuple[str, str], list[AlgebraicUnknown]] = {}
for endpoint in effort_unknowns:
members_by_root.setdefault(find(endpoint), []).append(
effort_unknowns[endpoint]
)
anchors_by_root: dict[tuple[str, str], list[float]] = {}
anchors_by_root: dict[
tuple[str, str],
list[tuple[AlgebraicUnknown, 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
for equation_variable in equation.variables
if (
(endpoint := self._port_key(variable, "p"))
in pressure_unknowns
(endpoint := self._port_key(equation_variable, variable))
in effort_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):
unknown = effort_unknowns[endpoint]
target_value = unknown.read() - float(equation.value)
if not isfinite(target_value):
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)
anchors_by_root.setdefault(find(endpoint), []).append(
(unknown, target_value)
)
for root, members in members_by_root.items():
anchors = anchors_by_root.get(root, [])
return tuple(
EffortEqualityGroup(
variable=variable,
members=tuple(members),
anchors=tuple(anchors_by_root.get(root, ())),
)
for root, members in members_by_root.items()
)
def _seed_equal_effort(self, variable: str) -> None:
for group in self._effort_equality_groups(variable):
members = group.members
anchors = group.anchors
if anchors:
pressure_scale = max([abs(value) for value in anchors] + [1.0])
if max(anchors) - min(anchors) > 1.0e-9 * pressure_scale:
# Keep each state-owned port current even when an invalid model
# has conflicting anchors in one equality group.
for unknown, target_value in anchors:
unknown.write(target_value)
anchor_values = [value for _unknown, value in anchors]
effort_scale = max([abs(value) for value in anchor_values] + [1.0])
if max(anchor_values) - min(anchor_values) > 1.0e-9 * effort_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)
target_value = sum(anchor_values) / len(anchor_values)
for unknown in members:
unknown.write(target_value)
continue
positive_seed = next(
(
pressure_unknowns[endpoint].read()
for endpoint in members
if pressure_unknowns[endpoint].read() > 0.0
),
None,
if variable == "p":
seed = next(
(
unknown.read()
for unknown in members
if unknown.read() > 0.0
),
None,
)
if seed is None:
continue
else:
seed = members[0].read()
for unknown in members:
if variable != "p" or unknown.read() <= 0.0:
unknown.write(seed)
def _connected_flow_unknown(
self,
component_name: str,
port_name: str,
variable: str,
) -> AlgebraicUnknown | None:
endpoint_key = (component_name, port_name)
for connection in self.network.connections:
if connection.kind != "physical":
continue
if connection.endpoint_a.key == endpoint_key:
other = connection.endpoint_b
elif connection.endpoint_b.key == endpoint_key:
other = connection.endpoint_a
else:
continue
return self._unknowns_by_id.get(
f"{other.component}.{other.port}.{variable}"
)
if positive_seed is None:
return None
@staticmethod
def _bisect_contact_root(
value_at,
lower: float,
upper: float,
target: float,
) -> float | None:
lower_value = float(value_at(lower)) - target
upper_value = float(value_at(upper)) - target
tolerance = 1.0e-13 * max(abs(target), 1.0)
if abs(lower_value) <= tolerance:
return lower
if abs(upper_value) <= tolerance:
return upper
if not isfinite(lower_value) or not isfinite(upper_value):
return None
if (lower_value < 0.0) == (upper_value < 0.0):
return None
for _iteration in range(100):
middle = 0.5 * (lower + upper)
middle_value = float(value_at(middle)) - target
if abs(middle_value) <= tolerance:
return middle
if (lower_value < 0.0) == (middle_value < 0.0):
lower = middle
lower_value = middle_value
else:
upper = middle
upper_value = middle_value
return 0.5 * (lower + upper)
def _contact_penetration_for_force(
self,
component,
requested_force: float,
current_penetration: float,
) -> float | None:
"""Invert one LSTP force law and select the root nearest its current state."""
if not isfinite(requested_force):
return None
option = int(getattr(component, "discContactOption", 2.0))
if option != 1:
requested_force = max(requested_force, 0.0)
stiffness = max(float(getattr(component, "kcont", 0.0)), 0.0)
damping = max(float(getattr(component, "rcont", 0.0)), 0.0)
damping_length = float(getattr(component, "Pdis", 0.0))
relative_velocity = float(getattr(component, "penetration_velocity"))
damping_term = damping * relative_velocity
current_penetration = (
max(float(current_penetration), 0.0)
if isfinite(current_penetration)
else 0.0
)
force_tolerance = 1.0e-12 * max(abs(requested_force), 1.0)
def raw_force(penetration: float) -> float:
if penetration <= 0.0:
return 0.0
damping_fraction = (
-expm1(-penetration / damping_length)
if damping_length > 0.0
else 1.0
)
return stiffness * penetration + damping_term * damping_fraction
def contact_force(penetration: float) -> float:
force = raw_force(penetration)
return force if option == 1 else max(force, 0.0)
candidates: list[float] = []
def add_candidate(penetration: float | None) -> None:
if penetration is None or not isfinite(penetration) or penetration < 0.0:
return
if abs(contact_force(penetration) - requested_force) > force_tolerance:
return
if not any(
abs(penetration - candidate)
<= 1.0e-12 * max(abs(penetration), abs(candidate), 1.0e-18)
for candidate in candidates
):
candidates.append(penetration)
add_candidate(current_penetration)
add_candidate(0.0)
if option != 1 and requested_force == 0.0:
return min(
candidates or [0.0],
key=lambda penetration: abs(penetration - current_penetration),
)
if damping_length <= 0.0:
if stiffness > 0.0:
penetration = (requested_force - damping_term) / stiffness
if penetration > 0.0:
add_candidate(penetration)
elif abs(requested_force - damping_term) <= force_tolerance:
add_candidate(max(current_penetration, 1.0e-18))
elif stiffness > 0.0:
critical_penetration: float | None = None
if damping_term < -stiffness * damping_length:
critical_penetration = damping_length * log(
-damping_term / (stiffness * damping_length)
)
add_candidate(critical_penetration)
upper = max(
current_penetration,
damping_length,
abs(requested_force) / stiffness,
critical_penetration or 0.0,
1.0e-18,
)
for _iteration in range(100):
upper_value = raw_force(upper)
if isfinite(upper_value) and upper_value >= requested_force:
break
upper *= 2.0
else:
upper = float("nan")
if isfinite(upper):
if critical_penetration is not None:
add_candidate(
self._bisect_contact_root(
raw_force,
0.0,
critical_penetration,
requested_force,
)
)
add_candidate(
self._bisect_contact_root(
raw_force,
critical_penetration,
upper,
requested_force,
)
)
else:
add_candidate(
self._bisect_contact_root(
raw_force,
0.0,
upper,
requested_force,
)
)
elif damping_term != 0.0:
upper = max(current_penetration, damping_length, 1.0e-18)
for _iteration in range(100):
upper_value = raw_force(upper)
crossed = (
upper_value >= requested_force
if damping_term > 0.0
else upper_value <= requested_force
)
if isfinite(upper_value) and crossed:
add_candidate(
self._bisect_contact_root(
raw_force,
0.0,
upper,
requested_force,
)
)
break
upper *= 2.0
if not candidates:
return None
return min(
candidates,
key=lambda penetration: abs(penetration - current_penetration),
)
def _apply_unilateral_contact_binding(
self,
binding: UnilateralContactBinding,
) -> bool:
component = binding.component
requested_force = binding.force_sign * binding.neighbor_force.read()
if int(getattr(component, "discContactOption", 2.0)) != 1:
requested_force = max(requested_force, 0.0)
cached_penetration = getattr(component, "_causal_penetration", None)
penetration = self._contact_penetration_for_force(
component,
requested_force,
(
float(cached_penetration)
if cached_penetration is not None
else float(getattr(component, "penetration"))
),
)
if penetration is None:
component.clear_causal_contact()
return False
gap0 = float(getattr(component, "gap0", 0.0))
if binding.algebraic_port == 1:
target = component.port_2.x - gap0 - penetration
else:
target = component.port_1.x + gap0 + penetration
for unknown in binding.algebraic_group.members:
unknown.write(target)
component.set_causal_contact(
penetration=penetration,
force=requested_force,
)
return True
def _refresh_unilateral_contacts(
self,
bindings: tuple[UnilateralContactBinding, ...],
) -> None:
for binding in bindings:
self._apply_unilateral_contact_binding(binding)
def _seed_unilateral_contacts(
self,
) -> tuple[UnilateralContactBinding, ...]:
"""Create local eliminations for contacts with one algebraic coordinate."""
position_groups = {
unknown.id: group
for group in self._effort_equality_groups("x")
for unknown in group.members
}
bindings: list[UnilateralContactBinding] = []
bound_group_ids: set[int] = set()
for component in self.network.components.values():
if component.model_type != "amesim_lstp00a":
continue
for endpoint in members:
unknown = pressure_unknowns[endpoint]
if unknown.read() <= 0.0:
unknown.write(positive_seed)
first_neighbor = self._connected_flow_unknown(
component.name,
"port_1",
"f",
)
second_neighbor = self._connected_flow_unknown(
component.name,
"port_2",
"f",
)
first_group = position_groups.get(f"{component.name}.port_1.x")
second_group = position_groups.get(f"{component.name}.port_2.x")
if (
first_group is None
or second_group is None
or first_group is second_group
):
continue
if not first_group.anchors and first_neighbor is not None:
binding = UnilateralContactBinding(
component=component,
algebraic_group=first_group,
neighbor_force=first_neighbor,
algebraic_port=1,
force_sign=1.0,
)
elif not second_group.anchors and second_neighbor is not None:
binding = UnilateralContactBinding(
component=component,
algebraic_group=second_group,
neighbor_force=second_neighbor,
algebraic_port=2,
force_sign=-1.0,
)
else:
# With both coordinates state-owned, penetration is a dynamic
# result rather than an algebraic active-set choice.
continue
group_id = id(binding.algebraic_group)
if group_id in bound_group_ids:
# One relative contact law may eliminate a free coordinate.
# Any other contact sharing that coordinate must remain in the
# nonlinear system or the projections would overwrite each
# other and make root selection order-dependent.
continue
if self._apply_unilateral_contact_binding(binding):
bindings.append(binding)
bound_group_ids.add(group_id)
def _seed_explicit_mass_flows(self) -> None:
"""Initialize explicit ``m_flow - f(...)`` constitutive relations.
return tuple(bindings)
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.
def _solve_explicit_flow_unknowns(self) -> set[str]:
"""Directly evaluate explicit flow variables before nonlinear closure.
Component constitutive equations use the normalized residual form
``flow_unknown + remainder = 0`` whenever exactly one physical flow
variable is present. Solve those relations by substitution first,
then propagate the known values through component balances and physical
connectors. This covers pneumatic ``m_flow`` variables as well as
mechanical forces ``f`` such as ``FORC`` without asking the nonlinear
optimizer to discover values many orders of magnitude away from zero.
The remaining coupled equations still go through ``least_squares``;
these assignments provide both a consistent initial guess and the
nominal magnitudes used to scale that smaller nonlinear problem.
"""
seeded_ids: set[str] = set()
# Mechanical reaction balances can contain null-space forces. Reusing
# an arbitrary least-squares distribution from the preceding RHS call
# makes contact activation history-dependent, so choose deterministic
# zero tear values and rebuild the force chain from current signals,
# states, and pressure loads on every closure.
for unknown in self.unknowns:
if unknown.variable == "f":
unknown.write(0.0)
# First evaluate constitutive relations that expose one flow unknown
# with unit coefficient. Other variables in the equation (pressure,
# displacement, velocity, or a signal) have already been refreshed for
# the current state and time by the staged system closure.
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 = [
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"
and self._unknowns_by_id[variable].role == "flow"
]
if len(mass_flow_unknowns) != 1:
if len(flow_unknowns) != 1:
continue
unknown = mass_flow_unknowns[0]
target_flow = unknown.read() - float(equation.value)
if not isfinite(target_flow):
unknown = flow_unknowns[0]
if unknown.id in seeded_ids:
continue
unknown.write(target_flow)
target_value = unknown.read() - float(equation.value)
if not isfinite(target_value):
continue
unknown.write(target_value)
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":
# V1/correctness-first implementation: repeatedly solve any balance that
# now has exactly one unknown flow variable left. Rebuilding and
# rescanning the complete residual tuple after every assignment keeps
# propagation deterministic, but costs O(flow unknowns * equations) and
# can dominate long, stiff simulations. A production follow-up should
# compile the assignment/tear order from the static topology once and
# evaluate only each owning component or connection residual here.
while True:
propagated = False
for equation in self.network.pressure_flow_equation_residuals():
if equation.role != "flow" or equation.relation not in {
"constitutive",
"sumToZero",
}:
continue
mass_flow_unknowns = [
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"
and self._unknowns_by_id[variable].role == "flow"
]
if len(mass_flow_unknowns) != 2:
if not flow_unknowns:
continue
seeded = [
unknown for unknown in mass_flow_unknowns if unknown.id in seeded_ids
variable_names = {unknown.variable for unknown in flow_unknowns}
if len(variable_names) != 1:
continue
unseeded = [
unknown for unknown in flow_unknowns if unknown.id not in seeded_ids
]
if len(seeded) != 1:
if len(unseeded) != 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)
unknown = unseeded[0]
target_value = unknown.read() - float(equation.value)
if not isfinite(target_value):
continue
unknown.write(target_value)
seeded_ids.add(unknown.id)
propagated = True
break
if not propagated:
# Causalize one remaining free flow in an otherwise normalized
# linear balance. This is the algebraic equivalent of choosing
# a tear variable: the other free flows retain their current
# guesses and one dependent flow closes the equation exactly.
# It also gives rank-deficient rigid-body reaction balances a
# deterministic starting point before state reduction supplies
# their common acceleration.
for equation in self.network.pressure_flow_equation_residuals():
if equation.role != "flow" or equation.relation not in {
"constitutive",
"sumToZero",
}:
continue
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].role == "flow"
]
unseeded = [
unknown
for unknown in flow_unknowns
if unknown.id not in seeded_ids
]
if len(unseeded) <= 1:
continue
if len({unknown.variable for unknown in flow_unknowns}) != 1:
continue
unknown = unseeded[-1]
target_value = unknown.read() - float(equation.value)
if not isfinite(target_value):
continue
unknown.write(target_value)
seeded_ids.add(unknown.id)
propagated = True
break
if not propagated:
break
# 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)
return seeded_ids
def _scales(self) -> dict[str, float]:
pressure_scale = max(
@@ -356,11 +756,28 @@ class PressureFlowSolver:
"Topology-driven simulation requires SciPy; install requirements.txt."
) from exc
self._seed_equal_pressures()
self._seed_explicit_mass_flows()
for component in self.network.components.values():
clear_causal_contact = getattr(component, "clear_causal_contact", None)
if clear_causal_contact is not None:
clear_causal_contact()
self._seed_equal_efforts()
self._solve_explicit_flow_unknowns()
contact_bindings = self._seed_unilateral_contacts()
if contact_bindings:
self._solve_explicit_flow_unknowns()
self._refresh_unilateral_contacts(contact_bindings)
scales = self._scales()
pressure_scale = scales["p"]
flow_scale = scales["m_flow"]
unknown_scales = {
unknown.id: (
max(abs(unknown.read()), 1.0)
if unknown.variable == "f"
else scales.get(unknown.variable, max(abs(unknown.read()), 1.0))
)
for unknown in self.unknowns
}
positive_pressures = [
unknown.read()
for unknown in self.unknowns
@@ -373,15 +790,28 @@ class PressureFlowSolver:
)
def variable_scale(unknown: AlgebraicUnknown) -> float:
return scales.get(unknown.variable, max(abs(unknown.read()), 1.0))
return unknown_scales[unknown.id]
def equation_scale(equation) -> float:
seeded_equations = self.network.pressure_flow_equation_residuals()
def initial_equation_scale(equation) -> float:
variable_names = [
variable.rsplit(".", 1)[-1]
for variable in equation.variables
]
if equation.role == "flow":
return scales["f"] if "f" in variable_names else flow_scale
force_scales = [
unknown_scales[variable]
for variable in equation.variables
if variable in self._unknowns_by_id
and self._unknowns_by_id[variable].variable == "f"
]
if force_scales:
# Freeze force scaling per equation. A 1e17 N source must
# not hide an unrelated 40 N piston/contact imbalance in a
# different mechanical branch.
return max(force_scales + [abs(float(equation.value)), 1.0])
return flow_scale
if equation.role == "effort":
if "x" in variable_names:
return scales["x"]
@@ -390,7 +820,14 @@ 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()
equation_scales = {
equation.id: initial_equation_scale(equation)
for equation in seeded_equations
}
def equation_scale(equation) -> float:
return equation_scales.get(equation.id, initial_equation_scale(equation))
seeded_scaled = [
abs(equation.value / equation_scale(equation))
for equation in seeded_equations
@@ -424,6 +861,12 @@ class PressureFlowSolver:
self.last_diagnostics = diagnostics
return diagnostics
# A causal contact retains its small relative penetration around the
# current absolute port coordinates. Keep that local coordinate during
# nonlinear fallback: the contact law remains responsive to optimizer
# increments, while a sub-ULP penetration is not lost by subtracting two
# large absolute displacements.
x0 = np.asarray(
[
(
@@ -450,6 +893,7 @@ class PressureFlowSolver:
def scaled_residuals(values):
assign(values)
self._refresh_unilateral_contacts(contact_bindings)
equations = self.network.pressure_flow_equation_residuals()
return np.asarray(
[
@@ -470,6 +914,7 @@ class PressureFlowSolver:
max_nfev=self.max_evaluations,
)
assign(result.x)
self._refresh_unilateral_contacts(contact_bindings)
equations = self.network.pressure_flow_equation_residuals()
scaled = [
abs(
+627
View File
@@ -0,0 +1,627 @@
from __future__ import annotations
from dataclasses import dataclass
from typing import Callable, Literal, Mapping, Sequence
from app.simulation.components.amesim.mechanical.translational import (
AmesimMecmas21,
)
from app.simulation.core.base import DynamicComponent
from app.simulation.solvers.solver import StateTransition
from app.simulation.systems.network import SimulationNetwork
ConstraintMode = Literal["uninitialized", "free", "lower", "upper"]
DenseState = Callable[[float], Sequence[float]]
@dataclass
class MechanicalConstraintGroup:
"""MECMAS21 inertias that share one rigid translational coordinate."""
components: tuple[AmesimMecmas21, ...]
mode: ConstraintMode = "uninitialized"
@property
def representative(self) -> AmesimMecmas21:
return self.components[0]
@property
def total_mass(self) -> float:
return sum(component.mass for component in self.components)
@property
def ideal_components(self) -> tuple[AmesimMecmas21, ...]:
return tuple(
component
for component in self.components
if component.uses_ideal_endstops
)
@property
def discrete_endstop_components(self) -> tuple[AmesimMecmas21, ...]:
return tuple(
component
for component in self.components
if int(component.stoptype) in {1, 3}
)
@property
def lower_bound(self) -> float | None:
components = self.discrete_endstop_components
return max((component.xmin for component in components), default=None)
@property
def upper_bound(self) -> float | None:
components = self.discrete_endstop_components
return min((component.xmax for component in components), default=None)
@staticmethod
def _boundary_tolerance(bound: float) -> float:
return 1.0e-12 * max(abs(bound), 1.0)
def reset_mode(self) -> None:
self.mode = "uninitialized"
def release(self) -> None:
self.mode = "free"
def synchronize_state(self) -> list[float]:
reference = self.representative
velocity_scale = max(
[abs(component.v) for component in self.components] + [1.0]
)
position_scale = max(
[abs(component.x) for component in self.components] + [1.0]
)
if any(
abs(component.v - reference.v) > 1.0e-10 * velocity_scale
or abs(component.x - reference.x) > 1.0e-10 * position_scale
for component in self.components[1:]
):
names = ", ".join(component.name for component in self.components)
raise ValueError(
"Rigidly connected MECMAS21 components must have consistent "
f"initial x/v states: {names}."
)
lower = self.lower_bound
upper = self.upper_bound
names = ", ".join(component.name for component in self.components)
if lower is not None and upper is not None and lower > upper:
raise ValueError(
"Rigidly connected MECMAS21 components have incompatible discrete "
f"endstop limits: {names}."
)
position = reference.x
below_lower = (
lower is not None
and position < lower - self._boundary_tolerance(lower)
)
above_upper = (
upper is not None
and position > upper + self._boundary_tolerance(upper)
)
if below_lower or above_upper:
raise ValueError(
f"Initial MECMAS21 position {position:g} is outside the discrete "
f"endstop limits for: {names}."
)
if lower is not None and position < lower:
position = lower
if upper is not None and position > upper:
position = upper
state = [reference.v, position]
self.set_state_vector(state)
return state
def set_state_vector(self, values: Sequence[float]) -> None:
state = [float(value) for value in values]
for component in self.components:
component.set_state_vector(state)
def total_unconstrained_force(self) -> float:
return sum(
component.mass * component.unconstrained_acceleration()
for component in self.components
)
def _static_endstop_side(self, total_force: float) -> str | None:
position = self.representative.x
velocity = self.representative.v
lower = self.lower_bound
upper = self.upper_bound
# MECMAS21's dvel is the friction stick threshold. Its discrete
# endstops release by motion direction; velocity away from a stop is free.
if (
lower is not None
and position <= lower + self._boundary_tolerance(lower)
and velocity <= 0.0
and total_force <= 0.0
):
return "lower"
if (
upper is not None
and position >= upper - self._boundary_tolerance(upper)
and velocity >= 0.0
and total_force >= 0.0
):
return "upper"
return None
def lock(self, side: Literal["lower", "upper"]) -> None:
self.mode = side
def impact_velocity(
self,
side: Literal["lower", "upper"],
incoming_velocity: float,
) -> float:
"""Return the post-impact velocity for the active group boundary."""
bound = self.lower_bound if side == "lower" else self.upper_bound
if bound is None:
return float(incoming_velocity)
parameter_name = "xmin" if side == "lower" else "xmax"
active_components = tuple(
component
for component in self.discrete_endstop_components
if abs(float(getattr(component, parameter_name)) - bound)
<= self._boundary_tolerance(bound)
)
if any(int(component.stoptype) == 1 for component in active_components):
return 0.0
restitution_components = tuple(
component
for component in active_components
if int(component.stoptype) == 3
)
speed = abs(float(incoming_velocity))
threshold = max(
(component.restdvel for component in restitution_components),
default=0.0,
)
if speed <= threshold:
return 0.0
# A rigid group cannot satisfy two different simultaneous rebounds;
# use the most dissipative active stop after plastic priority.
restitution = min(
(component.restcoeff for component in restitution_components),
default=0.0,
)
outgoing_speed = restitution * speed
return outgoing_speed if side == "lower" else -outgoing_speed
def update_acceleration(self) -> float:
"""Resolve the current ideal constraint without committing event mode.
ODE solvers may evaluate rejected or out-of-order trial states. The
derivative calculation therefore cannot change ``mode``; only an
accepted state transition may commit a discrete impact mode.
"""
total_force = self.total_unconstrained_force()
if self._static_endstop_side(total_force) is not None:
for component in self.components:
component.set_constraint_motion(
0.0,
velocity=0.0,
)
return 0.0
acceleration = total_force / self.total_mass
for component in self.components:
component.set_constraint_motion(acceleration)
return acceleration
StateEntry = DynamicComponent | MechanicalConstraintGroup
class MechanicalStateReducer:
"""V1 rigid-inertia reduction and event-driven discrete-endstop handling.
Rigid mechanical effort relations are causalized into one ``[v, x]`` ODE
coordinate per connected mass group. ``MECMAS21 stoptype=1`` applies a
plastic impact, while ``stoptype=3`` applies its restitution coefficient
above the configured velocity threshold.
"""
def __init__(
self,
network: SimulationNetwork,
dynamic_components: list[DynamicComponent],
) -> None:
self.network = network
self.dynamic_components = dynamic_components
self.groups = self._build_groups()
self._group_by_component = {
component.name: group
for group in self.groups
for component in group.components
}
self.state_entries = self._build_state_entries()
self._group_state_offsets = self._build_group_state_offsets()
@staticmethod
def _port_key(
variable: str,
expected_variable: str,
) -> tuple[str, str] | None:
try:
component, port, variable_name = variable.rsplit(".", 2)
except ValueError:
return None
if variable_name != expected_variable:
return None
return component, port
def _build_groups(self) -> tuple[MechanicalConstraintGroup, ...]:
mechanical_ports = {
(component.name, definition.name)
for component in self.network.components.values()
for definition in component.port_definitions
if definition.kind == "physical" and definition.domain == "mechanical"
}
parents = {
variable: {key: key for key in mechanical_ports}
for variable in ("x", "v")
}
def find(variable: str, key: tuple[str, str]) -> tuple[str, str]:
parent = parents[variable]
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(
variable: str,
first: tuple[str, str],
second: tuple[str, str],
) -> None:
first_root = find(variable, first)
second_root = find(variable, second)
if first_root != second_root:
parents[variable][second_root] = first_root
for connection in self.network.connections:
first = connection.endpoint_a.key
second = connection.endpoint_b.key
if first in mechanical_ports and second in mechanical_ports:
for variable in ("x", "v"):
union(variable, first, second)
for component in self.network.components.values():
for equation in component.pressure_flow_equation_residuals():
if equation.relation != "equal" or equation.role != "effort":
continue
for variable in ("x", "v"):
endpoints = [
endpoint
for equation_variable in equation.variables
if (
(endpoint := self._port_key(equation_variable, variable))
in mechanical_ports
)
]
for endpoint in endpoints[1:]:
union(variable, endpoints[0], endpoint)
masses = [
component
for component in self.dynamic_components
if isinstance(component, AmesimMecmas21)
]
for component in masses:
ports = [
(component.name, definition.name)
for definition in component.port_definitions
if definition.kind == "physical"
and definition.domain == "mechanical"
]
for port in ports[1:]:
for variable in ("x", "v"):
union(variable, ports[0], port)
masses_by_roots: dict[
tuple[tuple[str, str], tuple[str, str]],
list[AmesimMecmas21],
] = {}
for component in masses:
first_port = next(
(component.name, definition.name)
for definition in component.port_definitions
if definition.kind == "physical"
and definition.domain == "mechanical"
)
roots = (find("x", first_port), find("v", first_port))
masses_by_roots.setdefault(roots, []).append(component)
return tuple(
MechanicalConstraintGroup(tuple(components))
for components in masses_by_roots.values()
)
def _build_state_entries(self) -> tuple[StateEntry, ...]:
entries: list[StateEntry] = []
for component in self.dynamic_components:
group = self._group_by_component.get(component.name)
if group is None:
entries.append(component)
elif group.representative is component:
entries.append(group)
return tuple(entries)
def _build_group_state_offsets(self) -> dict[int, int]:
offsets: dict[int, int] = {}
cursor = 0
for entry in self.state_entries:
if isinstance(entry, MechanicalConstraintGroup):
offsets[id(entry)] = cursor
cursor += 2
else:
cursor += entry.state_size
return offsets
@property
def has_state_events(self) -> bool:
return any(group.discrete_endstop_components for group in self.groups)
def reset_constraint_modes(self) -> None:
for group in self.groups:
group.reset_mode()
def initial_state_vector(self) -> list[float]:
self.reset_constraint_modes()
values: list[float] = []
for entry in self.state_entries:
if isinstance(entry, MechanicalConstraintGroup):
values.extend(entry.synchronize_state())
else:
values.extend(entry.get_state_vector())
return values
def apply_state_vector(self, values: list[float]) -> None:
cursor = 0
for entry in self.state_entries:
state_size = (
2 if isinstance(entry, MechanicalConstraintGroup) else entry.state_size
)
next_cursor = cursor + state_size
state = values[cursor:next_cursor]
if isinstance(entry, MechanicalConstraintGroup):
entry.set_state_vector(state)
else:
entry.set_state_vector(state)
cursor = next_cursor
if cursor != len(values):
raise ValueError("State vector length does not match reduced dynamic components.")
def update_constraint_accelerations(self) -> None:
for group in self.groups:
group.update_acceleration()
def state_derivatives(
self,
connected_h: Mapping[str, Mapping[str, float]],
) -> list[float]:
derivatives: list[float] = []
for entry in self.state_entries:
component = (
entry.representative
if isinstance(entry, MechanicalConstraintGroup)
else entry
)
derivatives.extend(
component.state_derivative_from_ports(connected_h[component.name])
)
return derivatives
@staticmethod
def _locate_crossing(
dense_state: DenseState,
state_index: int,
bound: float,
side: Literal["lower", "upper"],
start_time: float,
end_time: float,
) -> float:
lower_time = float(start_time)
upper_time = float(end_time)
for _iteration in range(60):
middle_time = 0.5 * (lower_time + upper_time)
position = float(dense_state(middle_time)[state_index])
crossed = position <= bound if side == "lower" else position >= bound
if crossed:
upper_time = middle_time
else:
lower_time = middle_time
return upper_time
@staticmethod
def _locate_turnaround(
dense_state: DenseState,
velocity_index: int,
side: Literal["lower", "upper"],
start_time: float,
end_time: float,
) -> float:
"""Locate the velocity reversal preceding a same-step re-impact."""
lower_time = float(start_time)
upper_time = float(end_time)
for _iteration in range(60):
middle_time = 0.5 * (lower_time + upper_time)
velocity = float(dense_state(middle_time)[velocity_index])
turned = velocity <= 0.0 if side == "lower" else velocity >= 0.0
if turned:
upper_time = middle_time
else:
lower_time = middle_time
return upper_time
def state_transition(
self,
previous_time: float,
previous_state: list[float],
current_time: float,
current_state: list[float],
dense_state: DenseState,
) -> StateTransition | None:
"""Return the earliest discrete-endstop impact in one accepted ODE step."""
candidates: list[
tuple[float, MechanicalConstraintGroup, Literal["lower", "upper"], float]
] = []
for group in self.groups:
if not group.discrete_endstop_components:
continue
velocity_index = self._group_state_offsets[id(group)]
position_index = velocity_index + 1
previous_velocity = float(previous_state[velocity_index])
current_velocity = float(current_state[velocity_index])
previous_position = float(previous_state[position_index])
current_position = float(current_state[position_index])
lower = group.lower_bound
upper = group.upper_bound
if (
lower is not None
and previous_position <= lower + group._boundary_tolerance(lower)
and previous_velocity < 0.0
):
candidates.append((previous_time, group, "lower", lower))
elif (
lower is not None
and previous_position > lower
and current_position <= lower
):
candidates.append(
(
self._locate_crossing(
dense_state,
position_index,
lower,
"lower",
previous_time,
current_time,
),
group,
"lower",
lower,
)
)
elif (
lower is not None
and previous_position <= lower
and previous_velocity > 0.0
and current_velocity < 0.0
and current_position <= lower
):
turnaround_time = self._locate_turnaround(
dense_state,
velocity_index,
"lower",
previous_time,
current_time,
)
candidates.append(
(
self._locate_crossing(
dense_state,
position_index,
lower,
"lower",
turnaround_time,
current_time,
),
group,
"lower",
lower,
)
)
if (
upper is not None
and previous_position >= upper - group._boundary_tolerance(upper)
and previous_velocity > 0.0
):
candidates.append((previous_time, group, "upper", upper))
elif (
upper is not None
and previous_position < upper
and current_position >= upper
):
candidates.append(
(
self._locate_crossing(
dense_state,
position_index,
upper,
"upper",
previous_time,
current_time,
),
group,
"upper",
upper,
)
)
elif (
upper is not None
and previous_position >= upper
and previous_velocity < 0.0
and current_velocity > 0.0
and current_position >= upper
):
turnaround_time = self._locate_turnaround(
dense_state,
velocity_index,
"upper",
previous_time,
current_time,
)
candidates.append(
(
self._locate_crossing(
dense_state,
position_index,
upper,
"upper",
turnaround_time,
current_time,
),
group,
"upper",
upper,
)
)
if not candidates:
return None
event_time = min(candidate[0] for candidate in candidates)
event_state = [float(value) for value in dense_state(event_time)]
simultaneous_tolerance = 1.0e-12 * max(abs(event_time), 1.0)
for candidate_time, group, side, bound in candidates:
if abs(candidate_time - event_time) > simultaneous_tolerance:
continue
velocity_index = self._group_state_offsets[id(group)]
event_state[velocity_index] = group.impact_velocity(
side,
event_state[velocity_index],
)
event_state[velocity_index + 1] = bound
if event_state[velocity_index] == 0.0:
group.lock(side)
else:
group.release()
return StateTransition(time=event_time, state=event_state)
+423 -31
View File
@@ -8,6 +8,23 @@ from typing import Callable, Literal, Sequence
CancellationCheck = Callable[[], bool]
AcceptedStepCallback = Callable[[float], None]
IntegrationStatus = Literal["completed", "cancelled", "failed"]
DenseState = Callable[[float], list[float]]
@dataclass(frozen=True)
class StateTransition:
"""A state reset located inside an accepted integration step."""
time: float
state: list[float]
StateTransitionHandler = Callable[
[float, list[float], float, list[float], DenseState],
StateTransition | None,
]
_MAX_STATE_TRANSITIONS_AT_SAME_TIME = 64
class _IntegrationCancelled(Exception):
@@ -45,13 +62,123 @@ def _append_solution_sample(
time: float,
state: list[float],
) -> None:
if times and time <= times[-1] + 1e-12:
time = float(time)
if times and time <= times[-1]:
return
times.append(float(time))
times.append(time)
for index, value in enumerate(state):
states[index].append(float(value))
def _append_or_replace_solution_sample(
times: list[float],
states: list[list[float]],
time: float,
state: list[float],
) -> None:
"""Store a reset state even when its event time was already sampled."""
time = float(time)
if times and time == times[-1]:
times[-1] = time
for index, value in enumerate(state):
states[index][-1] = float(value)
return
_append_solution_sample(times, states, time, state)
def _normalize_state_transition(
transition: StateTransition,
before_time: float,
after_time: float,
state_size: int,
) -> StateTransition:
"""Validate and normalize a transition returned for an accepted step."""
if not isinstance(transition, StateTransition):
raise TypeError(
"State transition handlers must return StateTransition or None."
)
transition_time = float(transition.time)
if not math.isfinite(transition_time):
raise ValueError("State transition times must be finite numbers.")
tolerance = 16.0 * max(
math.ulp(before_time),
math.ulp(after_time),
math.ulp(transition_time),
)
if (
transition_time < before_time - tolerance
or transition_time > after_time + tolerance
):
raise ValueError(
"State transition time must lie inside the accepted integration step."
)
transition_time = min(max(transition_time, before_time), after_time)
transition_state = [float(value) for value in transition.state]
if len(transition_state) != state_size:
raise ValueError(
"State transition reset state must have the same size as the ODE state."
)
if not all(math.isfinite(value) for value in transition_state):
raise ValueError("State transition reset states must contain finite numbers.")
return StateTransition(time=transition_time, state=transition_state)
def _is_repeated_state_transition(
transition: StateTransition,
last_transition: StateTransition | None,
) -> bool:
"""Suppress only the exact reset that was just applied.
A second reset at the same instant is meaningful when it produces a
different state (for example, two constraints becoming active together).
"""
return (
last_transition is not None
and transition.time == last_transition.time
and transition.state == last_transition.state
)
def _next_same_time_transition_count(
transition: StateTransition,
last_transition: StateTransition | None,
previous_count: int,
) -> int:
count = (
previous_count + 1
if last_transition is not None
and transition.time == last_transition.time
else 1
)
if count > _MAX_STATE_TRANSITIONS_AT_SAME_TIME:
raise RuntimeError(
"State transition handler exceeded "
f"{_MAX_STATE_TRANSITIONS_AT_SAME_TIME} chained resets at the same time."
)
return count
def _align_transition_with_exact_endpoint(
transition: StateTransition,
requested_time: float,
exact_endpoint: float | None,
) -> StateTransition:
"""Keep an event reported at a breakpoint on that exact public timestamp."""
if exact_endpoint is not None and requested_time == exact_endpoint:
return StateTransition(
time=float(exact_endpoint),
state=list(transition.state),
)
return transition
def _normalize_breakpoints(
config: SolveIVPConfig,
breakpoints: Sequence[float] | None,
@@ -86,6 +213,7 @@ def _runge_kutta_4(
t_eval: list[float] | None,
cancel_check: CancellationCheck | None = None,
accepted_step_callback: AcceptedStepCallback | None = None,
state_transition_handler: StateTransitionHandler | None = None,
) -> ODESolution:
if t_eval is None:
point_count = max(
@@ -102,10 +230,22 @@ def _runge_kutta_4(
status: IntegrationStatus = "completed"
message = "Integrated with built-in RK4 fallback because SciPy is unavailable."
error: Exception | None = None
last_transition: StateTransition | None = None
same_time_transition_count = 0
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)
try:
for target_time in t_eval[1:]:
while current_time < target_time - 1e-15:
while current_time < target_time:
if cancel_check is not None and cancel_check():
raise _IntegrationCancelled
dt = min(config.max_step, target_time - current_time)
@@ -113,13 +253,63 @@ def _runge_kutta_4(
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 = [
next_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)
next_time = current_time + dt
transition: StateTransition | None = None
if state_transition_handler is not None:
step_start = current_time
step_state = list(state)
def dense_state(time: float) -> list[float]:
fraction = (float(time) - step_start) / (next_time - step_start)
return [
before + fraction * (after - before)
for before, after in zip(step_state, next_state)
]
candidate = state_transition_handler(
step_start,
list(step_state),
next_time,
list(next_state),
dense_state,
)
if candidate is not None:
candidate = _normalize_state_transition(
candidate,
step_start,
next_time,
len(state),
)
if not _is_repeated_state_transition(
candidate,
last_transition,
):
transition = candidate
if transition is not None:
same_time_transition_count = _next_same_time_transition_count(
transition,
last_transition,
same_time_transition_count,
)
current_time = transition.time
state = list(transition.state)
last_transition = transition
_append_or_replace_solution_sample(
times,
states,
current_time,
state,
)
else:
current_time = next_time
state = next_state
report_step(current_time)
_append_solution_sample(times, states, target_time, state)
except _IntegrationCancelled:
@@ -150,6 +340,7 @@ def _runge_kutta_4_segmented(
breakpoints: Sequence[float],
cancel_check: CancellationCheck | None = None,
accepted_step_callback: AcceptedStepCallback | None = None,
state_transition_handler: StateTransitionHandler | None = None,
) -> ODESolution:
"""RK4 fallback that never evaluates a pre-breakpoint step at the breakpoint."""
@@ -172,13 +363,15 @@ def _runge_kutta_4_segmented(
sample_index = 0
while (
sample_index < len(sample_times)
and sample_times[sample_index] <= config.t_start + 1e-12
and sample_times[sample_index] <= config.t_start
):
sample_index += 1
status: IntegrationStatus = "completed"
message = "Integrated with built-in RK4 fallback because SciPy is unavailable."
error: Exception | None = None
last_transition: StateTransition | None = None
same_time_transition_count = 0
last_reported_step: float | None = None
def report_step(time: float) -> None:
@@ -193,8 +386,8 @@ def _runge_kutta_4_segmented(
def advance_to(
target_time: float, reported_terminal_time: float | None = None
) -> None:
nonlocal current_time, state
while current_time < target_time - 1e-15:
nonlocal current_time, last_transition, same_time_transition_count, state
while current_time < target_time:
if cancel_check is not None and cancel_check():
raise _IntegrationCancelled
dt = min(config.max_step, target_time - current_time)
@@ -208,15 +401,72 @@ def _runge_kutta_4_segmented(
_vector_add(state, k2, 0.5 * dt),
)
k4 = rhs(current_time + dt, _vector_add(state, k3, dt))
state = [
next_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
next_time = current_time + dt
transition: StateTransition | None = None
if state_transition_handler is not None:
step_start = current_time
step_state = list(state)
def dense_state(time: float) -> list[float]:
fraction = (float(time) - step_start) / (next_time - step_start)
return [
before + fraction * (after - before)
for before, after in zip(step_state, next_state)
]
candidate = state_transition_handler(
step_start,
list(step_state),
next_time,
list(next_state),
dense_state,
)
if candidate is not None:
requested_time = float(candidate.time)
candidate = _normalize_state_transition(
candidate,
step_start,
next_time,
len(state),
)
candidate = _align_transition_with_exact_endpoint(
candidate,
requested_time,
reported_terminal_time,
)
if not _is_repeated_state_transition(
candidate,
last_transition,
):
transition = candidate
if transition is not None:
same_time_transition_count = _next_same_time_transition_count(
transition,
last_transition,
same_time_transition_count,
)
current_time = transition.time
state = list(transition.state)
last_transition = transition
_append_or_replace_solution_sample(
times,
states,
current_time,
state,
)
else:
current_time = next_time
state = next_state
report_time = current_time
if (
reported_terminal_time is not None
and current_time >= target_time - 1e-15
and current_time >= target_time
):
report_time = reported_terminal_time
report_step(report_time)
@@ -281,7 +531,16 @@ def _integrate_scipy_stepwise(
cancel_check: CancellationCheck,
accepted_step_callback: AcceptedStepCallback | None,
breakpoints: Sequence[float] = (),
state_transition_handler: StateTransitionHandler | None = None,
) -> ODESolution:
"""Initial stepwise integration path for breakpoints and state resets.
Known V1 limitation: an adaptive solver can evaluate a trial state outside
the algebraic or thermodynamic model domain. Such an RHS exception still
aborts the run here; recoverable trial failures are not yet restored to the
last accepted state and retried with a smaller step. This is not specific
to BDF, although implicit Newton/Jacobian probes make it especially visible.
"""
import numpy as np
from scipy.integrate import BDF, DOP853, LSODA, RK23, RK45, Radau
@@ -305,7 +564,7 @@ def _integrate_scipy_stepwise(
sample_index = 0
while (
sample_index < len(sample_times)
and sample_times[sample_index] <= config.t_start + 1e-12
and sample_times[sample_index] <= config.t_start
):
sample_index += 1
@@ -317,8 +576,18 @@ def _integrate_scipy_stepwise(
status: IntegrationStatus = "completed"
message = "The solver successfully reached the end of the integration interval."
error: Exception | None = None
last_transition: StateTransition | None = None
same_time_transition_count = 0
integration_progressed = False
last_reported_step: float | None = None
def cancellation_message() -> str:
return (
"Simulation was stopped before reaching the requested end time."
if integration_progressed
else "Simulation was stopped before integration started."
)
def report_step(time: float) -> None:
nonlocal last_reported_step
if accepted_step_callback is None:
@@ -332,11 +601,7 @@ def _integrate_scipy_stepwise(
for segment_index, segment_end in enumerate(segment_ends):
if cancel_check():
status = "cancelled"
message = (
"Simulation was stopped before integration started."
if segment_index == 0
else "Simulation was stopped before reaching the requested end time."
)
message = cancellation_message()
break
is_breakpoint = segment_index < len(breakpoints)
@@ -345,7 +610,12 @@ def _integrate_scipy_stepwise(
)
has_integration_interval = integration_end > last_accepted_time
if has_integration_interval:
while has_integration_interval and last_accepted_time < integration_end:
if cancel_check():
status = "cancelled"
message = cancellation_message()
break
solver_options = {
"rtol": config.rtol,
"atol": config.atol,
@@ -367,11 +637,7 @@ def _integrate_scipy_stepwise(
)
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."
)
message = cancellation_message()
break
except Exception as exc:
status = "failed"
@@ -379,6 +645,7 @@ def _integrate_scipy_stepwise(
error = exc
break
restart_at_transition = False
while solver.status == "running":
if cancel_check():
status = "cancelled"
@@ -386,6 +653,9 @@ def _integrate_scipy_stepwise(
"Simulation was stopped before reaching the requested end time."
)
break
step_start_time = last_accepted_time
step_start_state = list(last_accepted_state)
try:
step_message = solver.step()
except _IntegrationCancelled:
@@ -400,20 +670,118 @@ def _integrate_scipy_stepwise(
error = exc
break
integration_progressed = True
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]
step_end_time = float(solver.t)
step_end_state = [float(value) for value in solver.y]
dense_output = (
solver.dense_output()
if sample_times or state_transition_handler is not None
else None
)
transition: StateTransition | None = None
if state_transition_handler is not None:
assert dense_output is not None
def dense_state(time: float) -> list[float]:
return [float(value) for value in dense_output(float(time))]
try:
candidate = state_transition_handler(
step_start_time,
list(step_start_state),
step_end_time,
list(step_end_state),
dense_state,
)
if candidate is not None:
requested_time = float(candidate.time)
candidate = _normalize_state_transition(
candidate,
step_start_time,
step_end_time,
len(last_accepted_state),
)
candidate = _align_transition_with_exact_endpoint(
candidate,
requested_time,
float(segment_end) if is_breakpoint else None,
)
if not _is_repeated_state_transition(
candidate,
last_transition,
):
transition = candidate
except Exception as exc:
status = "failed"
message = str(exc)
error = exc
break
if transition is not None:
try:
same_time_transition_count = (
_next_same_time_transition_count(
transition,
last_transition,
same_time_transition_count,
)
)
except Exception as exc:
status = "failed"
message = str(exc)
error = exc
break
while (
sample_index < len(sample_times)
and sample_times[sample_index] < transition.time
):
sample_time = float(sample_times[sample_index])
assert dense_output is not None
sample_state = [
float(value) for value in dense_output(sample_time)
]
_append_solution_sample(
times,
states,
sample_time,
sample_state,
)
sample_index += 1
last_accepted_time = transition.time
last_accepted_state = list(transition.state)
last_transition = transition
_append_or_replace_solution_sample(
times,
states,
last_accepted_time,
last_accepted_state,
)
while (
sample_index < len(sample_times)
and sample_times[sample_index] <= last_accepted_time
):
sample_index += 1
report_step(last_accepted_time)
restart_at_transition = last_accepted_time < integration_end
break
last_accepted_time = step_end_time
last_accepted_state = step_end_state
reported_time = (
float(segment_end)
if is_breakpoint and solver.status == "finished"
else last_accepted_time
)
if sample_times:
dense_output = solver.dense_output()
assert dense_output is not None
while (
sample_index < len(sample_times)
and sample_times[sample_index] <= last_accepted_time
@@ -438,9 +806,12 @@ def _integrate_scipy_stepwise(
)
report_step(reported_time)
if status != "completed":
if status != "completed" or not restart_at_transition:
break
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
@@ -495,15 +866,29 @@ def integrate_ode(
cancel_check: CancellationCheck | None = None,
accepted_step_callback: AcceptedStepCallback | None = None,
breakpoints: Sequence[float] | None = None,
state_transition_handler: StateTransitionHandler | None = None,
):
"""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.
A state transition handler inspects every accepted step using its dense
interpolant. When it returns a transition, samples before the event retain
the pre-event trajectory, the reset state is stored at the event, and a fresh
solver continues from that state.
"""
if abs(config.t_stop - config.t_start) <= 1e-15:
if (
state_transition_handler is not None
and config.t_stop < config.t_start
):
raise ValueError(
"State transition handling does not support reverse integration."
)
if config.t_stop == config.t_start:
return ODESolution(
t=[float(config.t_start)],
y=[[value] for value in initial_state],
@@ -525,6 +910,7 @@ def integrate_ode(
normalized_breakpoints,
cancel_check,
accepted_step_callback,
state_transition_handler,
)
return _runge_kutta_4(
rhs,
@@ -533,9 +919,14 @@ def integrate_ode(
t_eval,
cancel_check,
accepted_step_callback,
state_transition_handler,
)
if cancel_check is not None or normalized_breakpoints:
if (
cancel_check is not None
or normalized_breakpoints
or state_transition_handler is not None
):
return _integrate_scipy_stepwise(
rhs,
initial_state,
@@ -544,6 +935,7 @@ def integrate_ode(
cancel_check or (lambda: False),
accepted_step_callback,
normalized_breakpoints,
state_transition_handler,
)
solve_options = {