验收四路模型并优化拓扑求解性能

This commit is contained in:
huojiarong committed 2026-08-12 11:57:42 +00:00
1 parent caca32a513
commit 456c29b3b6
20 files changed
+1085 -137

No files matched your search

+369 -103
View File
@@ -12,6 +12,7 @@ from app.simulation.components.amesim.flow.pipes import (
AmesimPnl0001,
AmesimPnl0002,
)
from app.simulation.core.equations import EquationResidual
from app.simulation.core.ports import PortState, VariableRole
from app.simulation.systems.network import SimulationNetwork
@@ -45,13 +46,45 @@ class AlgebraicUnknown:
class ExplicitFlowAssignment:
equation_id: str
unknown: AlgebraicUnknown
evaluate: Callable[[], float]
evaluate: Callable[[], float] | None
component: object | None = None
equation_index: int | None = None
@dataclass(frozen=True)
class ExplicitFlowStage:
assignments: tuple[ExplicitFlowAssignment, ...]
@dataclass(frozen=True)
class EffortAnchor:
unknown: AlgebraicUnknown
evaluate: Callable[[], float]
@dataclass(frozen=True)
class ConnectionEquationEvaluation:
template: EquationResidual
evaluate: Callable[[], float]
@dataclass(frozen=True)
class PnorPnl0001SeriesBinding:
orifice: AmesimPnor001
orifice_port: str
pipe: AmesimPnl0001
pipe_port: str
@dataclass(frozen=True)
class ClosedResistancePressureBinding:
component: object
port_name: str
neighbor: object
neighbor_port: str
pressure_source_port: str | None
@dataclass(frozen=True)
class EffortEqualityGroup:
variable: str
@@ -105,10 +138,43 @@ class PressureFlowSolver:
self.max_evaluations = max_evaluations
self.unknowns = self._build_unknowns()
self._unknowns_by_id = {unknown.id: unknown for unknown in self.unknowns}
self._unknowns_by_variable = {
variable: tuple(
unknown
for unknown in self.unknowns
if unknown.variable == variable
)
for variable in ("p", "m_flow", "x", "v", "f")
}
self._component_equation_owners = tuple(network.components.values())
self._estimated_flow_components = tuple(
component
for component in self._component_equation_owners
if hasattr(component, "K_eff")
)
self._causal_contact_components = tuple(
component
for component in self._component_equation_owners
if getattr(component, "clear_causal_contact", None) is not None
)
self._connection_equation_plan = tuple(
ConnectionEquationEvaluation(
template=equation,
evaluate=self._equation_value_reader(equation),
)
for equation in network.connection_equation_residuals()
)
self._effort_groups = {
variable: self._build_effort_equality_groups(variable)
for variable in ("p", "x", "v")
}
self._pnor_pnl0001_series_plan = (
self._build_pnor_pnl0001_series_plan()
)
self._closed_resistance_pressure_plan = (
self._build_closed_resistance_pressure_plan()
)
self._unilateral_contact_plan = self._build_unilateral_contact_plan()
self._explicit_flow_plan = self._build_explicit_flow_plan()
self.last_diagnostics: AlgebraicSolveDiagnostics | None = None
@@ -143,7 +209,10 @@ class PressureFlowSolver:
return None
return component_name, port_name
def _seed_equal_efforts(self) -> None:
def _seed_equal_efforts(
self,
variables: tuple[str, ...] = ("p", "x", "v"),
) -> None:
"""Lift state-owned efforts across their complete equality groups.
Dynamic components refresh their own ports before each closure, while
@@ -154,7 +223,19 @@ class PressureFlowSolver:
before evaluating explicit flow laws.
"""
for variable in ("p", "x", "v"):
self.propagate_equal_efforts(variables)
def propagate_equal_efforts(self, variables: tuple[str, ...]) -> None:
"""Propagate selected state-owned efforts without solving flows.
Piston geometry needs current mechanical ``x``/``v`` before swept
volume propagation, but pressure and flow equations can wait until the
connected chamber has refreshed that volume.
"""
unknown = sorted(set(variables) - set(self._effort_groups))
if unknown:
raise ValueError("Unsupported effort variables: " + ", ".join(unknown))
for variable in variables:
self._seed_equal_effort(variable)
def _build_effort_equality_groups(
@@ -538,10 +619,10 @@ class PressureFlowSolver:
for binding in bindings:
self._apply_unilateral_contact_binding(binding)
def _seed_unilateral_contacts(
def _build_unilateral_contact_plan(
self,
) -> tuple[UnilateralContactBinding, ...]:
"""Create local eliminations for contacts with one algebraic coordinate."""
"""Compile contacts that can eliminate one algebraic coordinate."""
position_groups = {
unknown.id: group
@@ -549,7 +630,6 @@ class PressureFlowSolver:
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
@@ -591,6 +671,18 @@ class PressureFlowSolver:
# With both coordinates state-owned, penetration is a dynamic
# result rather than an algebraic active-set choice.
continue
bindings.append(binding)
return tuple(bindings)
def _seed_unilateral_contacts(
self,
) -> tuple[UnilateralContactBinding, ...]:
"""Apply compiled local contact eliminations for the current state."""
bindings: list[UnilateralContactBinding] = []
bound_group_ids: set[int] = set()
for binding in self._unilateral_contact_plan:
group_id = id(binding.algebraic_group)
if group_id in bound_group_ids:
# One relative contact law may eliminate a free coordinate.
@@ -630,33 +722,87 @@ class PressureFlowSolver:
component = self.network.components[equation.owner_id]
equation_id = equation.id
def read_component_equation() -> float:
for current in component.pressure_flow_equation_residuals():
if current.id == equation_id:
return float(current.value)
equation_ids = tuple(
current.id
for current in component.pressure_flow_equation_residuals()
)
try:
equation_index = equation_ids.index(equation_id)
except ValueError as exc:
raise RuntimeError(
f"Compiled algebraic equation disappeared at runtime: {equation_id}."
)
) from exc
def read_component_equation() -> float:
current_equations = component.pressure_flow_equation_residuals()
if (
equation_index >= len(current_equations)
or current_equations[equation_index].id != equation_id
):
raise RuntimeError(
f"Compiled algebraic equation disappeared at runtime: {equation_id}."
)
return float(current_equations[equation_index].value)
return read_component_equation
def _pressure_flow_equation_residuals(self) -> tuple[EquationResidual, ...]:
"""Evaluate live values through a precompiled connector topology."""
component_residuals = tuple(
residual
for component in self._component_equation_owners
for residual in component.pressure_flow_equation_residuals()
)
connection_residuals = tuple(
EquationResidual(
id=item.template.id,
owner=item.template.owner,
owner_id=item.template.owner_id,
relation=item.template.relation,
variables=item.template.variables,
value=item.evaluate(),
role=item.template.role,
)
for item in self._connection_equation_plan
)
return component_residuals + connection_residuals
def _build_explicit_flow_plan(self) -> tuple[ExplicitFlowAssignment, ...]:
"""Compile the legacy deterministic flow assignment order once."""
assignments: list[ExplicitFlowAssignment] = []
def _explicit_flow_assignment(
self,
equation,
unknown: AlgebraicUnknown,
) -> ExplicitFlowAssignment:
if equation.owner == "connection":
return ExplicitFlowAssignment(
equation_id=equation.id,
unknown=unknown,
evaluate=self._equation_value_reader(equation),
)
component = self.network.components[equation.owner_id]
equations = component.pressure_flow_equation_residuals()
equation_ids = tuple(current.id for current in equations)
try:
equation_index = equation_ids.index(equation.id)
except ValueError as exc:
raise RuntimeError(
f"Compiled algebraic equation disappeared at runtime: {equation.id}."
) from exc
return ExplicitFlowAssignment(
equation_id=equation.id,
unknown=unknown,
evaluate=None,
component=component,
equation_index=equation_index,
)
def _build_explicit_flow_plan(self) -> tuple[ExplicitFlowStage, ...]:
"""Compile flow causalization into independent dependency stages."""
stages: list[ExplicitFlowStage] = []
seeded_ids: set[str] = set()
def append_assignment(equation, unknown: AlgebraicUnknown) -> None:
assignments.append(
ExplicitFlowAssignment(
equation_id=equation.id,
unknown=unknown,
evaluate=self._equation_value_reader(equation),
)
)
seeded_ids.add(unknown.id)
initial_assignments: list[ExplicitFlowAssignment] = []
for component in self.network.components.values():
for equation in component.pressure_flow_equation_residuals():
if equation.relation != "constitutive" or equation.role != "flow":
@@ -665,12 +811,19 @@ class PressureFlowSolver:
if len(flow_unknowns) != 1:
continue
unknown = flow_unknowns[0]
if unknown.id not in seeded_ids:
append_assignment(equation, unknown)
if unknown.id in seeded_ids:
continue
initial_assignments.append(
self._explicit_flow_assignment(equation, unknown)
)
seeded_ids.add(unknown.id)
if initial_assignments:
stages.append(ExplicitFlowStage(tuple(initial_assignments)))
equations = self.network.pressure_flow_equation_residuals()
equations = self._pressure_flow_equation_residuals()
while True:
propagated = False
stage_assignments: list[ExplicitFlowAssignment] = []
stage_unknown_ids: set[str] = set()
for equation in equations:
if equation.role != "flow" or equation.relation not in {
"constitutive",
@@ -689,12 +842,19 @@ class PressureFlowSolver:
)
if len(unseeded) != 1:
continue
append_assignment(equation, unseeded[0])
propagated = True
break
if propagated:
unknown = unseeded[0]
if unknown.id in stage_unknown_ids:
continue
stage_assignments.append(
self._explicit_flow_assignment(equation, unknown)
)
stage_unknown_ids.add(unknown.id)
if stage_assignments:
stages.append(ExplicitFlowStage(tuple(stage_assignments)))
seeded_ids.update(stage_unknown_ids)
continue
fallback_assignment: ExplicitFlowAssignment | None = None
for equation in equations:
if equation.role != "flow" or equation.relation not in {
"constitutive",
@@ -711,40 +871,92 @@ class PressureFlowSolver:
continue
if len({unknown.variable for unknown in flow_unknowns}) != 1:
continue
append_assignment(equation, unseeded[-1])
propagated = True
unknown = unseeded[-1]
fallback_assignment = self._explicit_flow_assignment(
equation,
unknown,
)
seeded_ids.add(unknown.id)
break
if not propagated:
if fallback_assignment is None:
break
stages.append(ExplicitFlowStage((fallback_assignment,)))
return tuple(assignments)
return tuple(stages)
def _solve_explicit_flow_unknowns(self) -> set[str]:
"""Execute the precompiled explicit flow/force causalization plan."""
@staticmethod
def _evaluate_explicit_flow_stage(
assignments: tuple[ExplicitFlowAssignment, ...],
) -> dict[str, float]:
values: dict[str, float] = {}
assignments_by_component: dict[object, list[ExplicitFlowAssignment]] = {}
for assignment in assignments:
if assignment.component is None:
assert assignment.evaluate is not None
values[assignment.equation_id] = assignment.evaluate()
continue
assignments_by_component.setdefault(assignment.component, []).append(
assignment
)
for component, component_assignments in assignments_by_component.items():
equations = component.pressure_flow_equation_residuals()
for assignment in component_assignments:
assert assignment.equation_index is not None
equation_index = assignment.equation_index
if (
equation_index >= len(equations)
or equations[equation_index].id != assignment.equation_id
):
raise RuntimeError(
"Compiled algebraic equation disappeared at runtime: "
f"{assignment.equation_id}."
)
values[assignment.equation_id] = float(
equations[equation_index].value
)
return values
def _solve_explicit_flow_unknowns(
self,
variables: tuple[str, ...] = ("f", "m_flow"),
) -> set[str]:
"""Execute staged flow/force assignments without repeated equations."""
selected = frozenset(variables)
for unknown in self.unknowns:
if unknown.variable in {"f", "m_flow"}:
if unknown.variable in selected:
unknown.write(0.0)
seeded_ids: set[str] = set()
for assignment in self._explicit_flow_plan:
target_value = assignment.unknown.read() - assignment.evaluate()
if not isfinite(target_value):
continue
assignment.unknown.write(target_value)
seeded_ids.add(assignment.unknown.id)
for stage in self._explicit_flow_plan:
assignments = tuple(
assignment
for assignment in stage.assignments
if assignment.unknown.variable in selected
)
values = self._evaluate_explicit_flow_stage(assignments)
targets = tuple(
(
assignment,
assignment.unknown.read() - values[assignment.equation_id],
)
for assignment in assignments
)
for assignment, target_value in targets:
if not isfinite(target_value):
continue
assignment.unknown.write(target_value)
seeded_ids.add(assignment.unknown.id)
return seeded_ids
def _seed_closed_resistance_pressures(self) -> None:
"""Seed a sealed resistance end at its zero-flow pressure.
A PNPL01 fixes flow, not pressure. Starting a dead-ended Darcy branch
with the plug-side pressure at the medium reference can otherwise put
the nonlinear solver on the singular square-root part of the inverse
flow law. At zero flow, these AMESim pipe resistances have exactly zero
pressure drop, which gives a deterministic and physically exact seed.
"""
def _build_closed_resistance_pressure_plan(
self,
) -> tuple[ClosedResistancePressureBinding, ...]:
"""Compile sealed resistance ends whose zero-flow pressure is known."""
bindings: list[ClosedResistancePressureBinding] = []
connected: dict[tuple[str, str], tuple[str, str]] = {}
for connection in self.network.connections:
if connection.kind != "physical" or connection.domain != "pneumatic":
@@ -764,20 +976,50 @@ class PressureFlowSolver:
if not isinstance(neighbor, AmesimPnpl01):
continue
if isinstance(component, AmesimPnl0002):
pressure = component.properties().p
pressure_source_port = None
elif isinstance(component, AmesimPnl0001):
if port_name != "port_1":
continue
pressure = component.properties().p
pressure_source_port = None
else:
other_port_name = "port_2" if port_name == "port_1" else "port_1"
pressure = component.get_port(other_port_name).p
component.get_port(port_name).p = pressure
neighbor.get_port(neighbor_key[1]).p = pressure
pressure_source_port = (
"port_2" if port_name == "port_1" else "port_1"
)
bindings.append(
ClosedResistancePressureBinding(
component=component,
port_name=port_name,
neighbor=neighbor,
neighbor_port=neighbor_key[1],
pressure_source_port=pressure_source_port,
)
)
return tuple(bindings)
def _seed_pnor_pnl0001_series_pressures(self) -> None:
"""Causalize the pressure between a PNOR001 and PNL0001 R port."""
def _seed_closed_resistance_pressures(self) -> None:
"""Seed a sealed resistance end at its zero-flow pressure.
A PNPL01 fixes flow, not pressure. Starting a dead-ended Darcy branch
with the plug-side pressure at the medium reference can otherwise put
the nonlinear solver on the singular square-root part of the inverse
flow law. At zero flow, these AMESim pipe resistances have exactly zero
pressure drop, which gives a deterministic and physically exact seed.
"""
for binding in self._closed_resistance_pressure_plan:
component = binding.component
pressure = (
component.properties().p
if binding.pressure_source_port is None
else component.get_port(binding.pressure_source_port).p
)
component.get_port(binding.port_name).p = pressure
binding.neighbor.get_port(binding.neighbor_port).p = pressure
def _build_pnor_pnl0001_series_plan(
self,
) -> tuple[PnorPnl0001SeriesBinding, ...]:
bindings: list[PnorPnl0001SeriesBinding] = []
for connection in self.network.connections:
first_endpoint, second_endpoint = connection.endpoints
first = self.network.components[first_endpoint.component]
@@ -792,6 +1034,26 @@ class PressureFlowSolver:
continue
if isinstance(pipe, AmesimPnl0002) or pipe_port != "port_1":
continue
bindings.append(
PnorPnl0001SeriesBinding(
orifice=orifice,
orifice_port=orifice_port,
pipe=pipe,
pipe_port=pipe_port,
)
)
return tuple(bindings)
def _seed_pnor_pnl0001_series_pressures(self) -> None:
"""Causalize the pressure between a PNOR001 and PNL0001 R port."""
from scipy.optimize import brentq
for binding in self._pnor_pnl0001_series_plan:
orifice = binding.orifice
orifice_port = binding.orifice_port
pipe = binding.pipe
pipe_port = binding.pipe_port
orifice_other = "port_2" if orifice_port == "port_1" else "port_1"
pressure_a = orifice.get_port(orifice_other).p
@@ -826,15 +1088,16 @@ class PressureFlowSolver:
elif (lower_value < 0.0) == (upper_value < 0.0):
continue
else:
for _iteration in range(64):
middle = 0.5 * (lower + upper)
middle_value = mismatch(middle)
if (middle_value < 0.0) == (lower_value < 0.0):
lower = middle
lower_value = middle_value
else:
upper = middle
pressure = 0.5 * (lower + upper)
pressure = float(
brentq(
mismatch,
lower,
upper,
xtol=1.0e-6,
rtol=1.0e-12,
maxiter=32,
)
)
orifice.get_port(orifice_port).p = pressure
pipe.get_port(pipe_port).p = pressure
@@ -842,22 +1105,20 @@ class PressureFlowSolver:
pressure_scale = max(
[
abs(unknown.read())
for unknown in self.unknowns
if unknown.variable == "p" and unknown.read() > 0.0
for unknown in self._unknowns_by_variable["p"]
if unknown.read() > 0.0
]
+ [1e5]
)
estimated_flows = [
abs(float(getattr(component, "K_eff"))) * sqrt(pressure_scale)
for component in self.network.components.values()
if hasattr(component, "K_eff")
for component in self._estimated_flow_components
]
mass_flow_scale = max(
estimated_flows
+ [
abs(unknown.read())
for unknown in self.unknowns
if unknown.variable == "m_flow"
for unknown in self._unknowns_by_variable["m_flow"]
]
+ [1e-3]
)
@@ -865,20 +1126,24 @@ class PressureFlowSolver:
"p": pressure_scale,
"m_flow": mass_flow_scale,
"x": max(
[abs(unknown.read()) for unknown in self.unknowns if unknown.variable == "x"]
[abs(unknown.read()) for unknown in self._unknowns_by_variable["x"]]
+ [1.0]
),
"v": max(
[abs(unknown.read()) for unknown in self.unknowns if unknown.variable == "v"]
[abs(unknown.read()) for unknown in self._unknowns_by_variable["v"]]
+ [1.0]
),
"f": max(
[abs(unknown.read()) for unknown in self.unknowns if unknown.variable == "f"]
[abs(unknown.read()) for unknown in self._unknowns_by_variable["f"]]
+ [1.0]
),
}
def solve(self) -> AlgebraicSolveDiagnostics:
def solve(
self,
*,
effort_variables: tuple[str, ...] = ("p", "x", "v"),
) -> AlgebraicSolveDiagnostics:
try:
import numpy as np
from scipy.optimize import least_squares
@@ -887,18 +1152,16 @@ class PressureFlowSolver:
"Topology-driven simulation requires SciPy; install requirements.txt."
) from exc
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()
for component in self._causal_contact_components:
component.clear_causal_contact()
self._seed_equal_efforts()
self._seed_equal_efforts(effort_variables)
self._seed_closed_resistance_pressures()
self._seed_pnor_pnl0001_series_pressures()
self._solve_explicit_flow_unknowns()
contact_bindings = self._seed_unilateral_contacts()
if contact_bindings:
self._solve_explicit_flow_unknowns()
self._solve_explicit_flow_unknowns(("f",))
self._refresh_unilateral_contacts(contact_bindings)
scales = self._scales()
pressure_scale = scales["p"]
@@ -911,21 +1174,10 @@ class PressureFlowSolver:
)
for unknown in self.unknowns
}
positive_pressures = [
unknown.read()
for unknown in self.unknowns
if unknown.variable == "p" and unknown.read() > 0.0
]
fallback_pressure = (
sum(positive_pressures) / len(positive_pressures)
if positive_pressures
else pressure_scale
)
def variable_scale(unknown: AlgebraicUnknown) -> float:
return unknown_scales[unknown.id]
seeded_equations = self.network.pressure_flow_equation_residuals()
seeded_equations = self._pressure_flow_equation_residuals()
def initial_equation_scale(equation) -> float:
variable_names = [
@@ -959,7 +1211,10 @@ class PressureFlowSolver:
}
def equation_scale(equation) -> float:
return equation_scales.get(equation.id, initial_equation_scale(equation))
cached = equation_scales.get(equation.id)
if cached is not None:
return cached
return initial_equation_scale(equation)
seeded_scaled = [
abs(equation.value / equation_scale(equation))
@@ -994,6 +1249,17 @@ class PressureFlowSolver:
self.last_diagnostics = diagnostics
return diagnostics
positive_pressures = [
unknown.read()
for unknown in self._unknowns_by_variable["p"]
if unknown.read() > 0.0
]
fallback_pressure = (
sum(positive_pressures) / len(positive_pressures)
if positive_pressures
else pressure_scale
)
# 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
@@ -1027,7 +1293,7 @@ class PressureFlowSolver:
def scaled_residuals(values):
assign(values)
self._refresh_unilateral_contacts(contact_bindings)
equations = self.network.pressure_flow_equation_residuals()
equations = self._pressure_flow_equation_residuals()
return np.asarray(
[
equation.value / equation_scale(equation)
@@ -1048,7 +1314,7 @@ class PressureFlowSolver:
)
assign(result.x)
self._refresh_unilateral_contacts(contact_bindings)
equations = self.network.pressure_flow_equation_residuals()
equations = self._pressure_flow_equation_residuals()
scaled = [
abs(
equation.value / equation_scale(equation)
+23 -2
View File
@@ -391,6 +391,27 @@ class MechanicalStateReducer:
def has_state_events(self) -> bool:
return any(group.discrete_endstop_components for group in self.groups)
def absolute_tolerances(
self,
default: float,
*,
mechanical: float = 1.0e-12,
) -> list[float]:
"""Return state-aligned tolerances with machine-scale mechanics.
A scalar ``1e-8`` absolute tolerance makes SciPy perturb a zero-valued
endstop position across the much smaller unilateral boundary band while
constructing finite-difference Jacobians. Mechanical coordinates need
a tighter floor; thermodynamic states retain the caller's tolerance.
"""
values: list[float] = []
for entry in self.state_entries:
if isinstance(entry, MechanicalConstraintGroup):
values.extend([min(default, mechanical)] * 2)
else:
values.extend([default] * entry.state_size)
return values
def reset_constraint_modes(self) -> None:
for group in self.groups:
group.reset_mode()
@@ -518,7 +539,7 @@ class MechanicalStateReducer:
candidates.append((previous_time, group, "lower", lower))
elif (
lower is not None
and previous_position > lower
and previous_position > lower + group._boundary_tolerance(lower)
and current_position <= lower
):
candidates.append(
@@ -573,7 +594,7 @@ class MechanicalStateReducer:
candidates.append((previous_time, group, "upper", upper))
elif (
upper is not None
and previous_position < upper
and previous_position < upper - group._boundary_tolerance(upper)
and current_position >= upper
):
candidates.append(
+43 -1
View File
@@ -39,7 +39,7 @@ class SolveIVPConfig:
t_stop: float = 20.0
method: str = "BDF"
rtol: float = 1e-6
atol: float = 1e-8
atol: float | Sequence[float] = 1e-8
max_step: float = 1e-3
first_step: float | None = None
@@ -88,6 +88,34 @@ def _append_or_replace_solution_sample(
return
_append_solution_sample(times, states, time, state)
def _project_nearby_pre_transition_sample(
times: list[float],
states: list[list[float]],
transition: StateTransition,
config: SolveIVPConfig,
) -> None:
"""Resolve a sample/event ordering that is below solver time precision.
An adaptive dense interpolant can place a discontinuous impact a few
nanoseconds after its analytically coincident output sample. Keep the
located event and restart time unchanged, but report that ambiguous sample
on the reset side of the discontinuity.
"""
if not times or not math.isfinite(config.max_step):
return
time_gap = float(transition.time) - times[-1]
tolerance = max(
64.0 * math.ulp(max(abs(float(transition.time)), 1.0)),
min(
abs(float(config.max_step) * float(config.rtol)),
1.0e-8,
),
)
if not 0.0 < time_gap <= tolerance:
return
for index, value in enumerate(transition.state):
states[index][-1] = float(value)
def _normalize_state_transition(
transition: StateTransition,
@@ -534,6 +562,7 @@ def _integrate_scipy_stepwise(
accepted_step_callback: AcceptedStepCallback | None,
breakpoints: Sequence[float] = (),
state_transition_handler: StateTransitionHandler | None = None,
jac_sparsity=None,
) -> ODESolution:
"""Initial stepwise integration path for breakpoints and state resets.
@@ -625,6 +654,8 @@ def _integrate_scipy_stepwise(
"atol": config.atol,
"max_step": segment_max_step,
}
if jac_sparsity is not None and config.method in {"BDF", "Radau"}:
solver_options["jac_sparsity"] = jac_sparsity
requested_first_step = (
0.1 * segment_max_step
if last_recoverable_error is not None
@@ -666,6 +697,7 @@ def _integrate_scipy_stepwise(
error = exc
break
restart_at_transition = False
restart_after_recoverable = False
while solver.status == "running":
@@ -806,6 +838,12 @@ def _integrate_scipy_stepwise(
)
sample_index += 1
_project_nearby_pre_transition_sample(
times,
states,
transition,
config,
)
last_accepted_time = transition.time
last_accepted_state = list(transition.state)
last_transition = transition
@@ -923,6 +961,7 @@ def integrate_ode(
accepted_step_callback: AcceptedStepCallback | None = None,
breakpoints: Sequence[float] | None = None,
state_transition_handler: StateTransitionHandler | None = None,
jac_sparsity=None,
):
"""Integrate an ODE, optionally restarting at equation discontinuities.
@@ -992,6 +1031,7 @@ def integrate_ode(
accepted_step_callback,
normalized_breakpoints,
state_transition_handler,
jac_sparsity,
)
solve_options = {
@@ -1006,4 +1046,6 @@ def integrate_ode(
}
if config.first_step is not None:
solve_options["first_step"] = config.first_step
if jac_sparsity is not None and config.method in {"BDF", "Radau"}:
solve_options["jac_sparsity"] = jac_sparsity
return solve_ivp(**solve_options)