Files
SystemSimulationApp/app/simulation/solvers/solver.py
T
lujingze b435daecf2 完善通用求解器回归与前端交互
- 引入因果坐标内核、热流体恢复和递进长时回归\n- 完善正交连线、线桥、视图保持与结果曲线缩放\n- 补充依赖约束、CI、测试基线和北京时间更新日志
2026-08-18 06:42:07 +00:00

1479 lines
54 KiB
Python

from __future__ import annotations
from app.simulation.core.errors import RecoverableTrialStateError
from app.simulation.performance import profile_phase
import math
from dataclasses import dataclass
from typing import Callable, Literal, Sequence
CancellationCheck = Callable[[], bool]
AcceptedStepCallback = Callable[[float], None]
IntegrationStatus = Literal["completed", "cancelled", "failed"]
DenseState = Callable[[float], list[float]]
JacobianCallable = Callable[[float, object], object]
@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
_MAX_RECOVERABLE_RETRIES = 16
_RECOVERABLE_RETRY_FACTOR = 0.5
class IntegrationCancelled(Exception):
"""Internal control-flow signal shared by RHS and Jacobian evaluation."""
pass
@dataclass(frozen=True)
class SolveIVPConfig:
t_start: float = 0.0
t_stop: float = 20.0
method: str = "BDF"
rtol: float = 1e-6
atol: float | Sequence[float] = 1e-8
max_step: float = 1e-3
first_step: float | None = None
@dataclass(frozen=True)
class RecoverableRetryDiagnostics:
"""One recoverable trial failure and the step cap chosen for its retry."""
phase: Literal["constructor", "step", "solver-status"]
attempted_step: float
reason: str
next_max_step: float | None = None
next_first_step: float | None = None
def as_dict(self) -> dict[str, object]:
result: dict[str, object] = {
"phase": self.phase,
"attemptedStep": self.attempted_step,
"reason": self.reason,
}
if self.next_max_step is not None:
result["nextMaxStep"] = self.next_max_step
if self.next_first_step is not None:
result["nextFirstStep"] = self.next_first_step
return result
@dataclass(frozen=True)
class SolverSegmentDiagnostics:
"""Work performed by implicit solver instances inside one event segment."""
start_time: float
requested_stop_time: float
simulated_until: float
nfev: int = 0
njev: int = 0
nlu: int = 0
accepted_step_count: int = 0
solver_start_count: int = 0
state_transition_count: int = 0
state_transition_times: tuple[float, ...] = ()
recoverable_retry_count: int = 0
jacobian_evaluation_count: int = 0
jacobian_full_build_count: int = 0
jacobian_secant_reuse_count: int = 0
jacobian_audit_failure_count: int = 0
finite_difference_rhs_evaluation_count: int = 0
jacobian_base_rhs_evaluation_count: int = 0
jacobian_jv_audit_rhs_evaluation_count: int = 0
exact_column_build_count: int = 0
exact_column_fallback_count: int = 0
jacobian_assembly_seconds: float = 0.0
recoverable_retries: tuple[RecoverableRetryDiagnostics, ...] = ()
def as_dict(self) -> dict[str, object]:
result: dict[str, object] = {
"startTime": self.start_time,
"requestedStopTime": self.requested_stop_time,
"simulatedUntil": self.simulated_until,
"nfev": self.nfev,
"njev": self.njev,
"nlu": self.nlu,
"acceptedStepCount": self.accepted_step_count,
"solverStartCount": self.solver_start_count,
"stateTransitionCount": self.state_transition_count,
"recoverableRetryCount": self.recoverable_retry_count,
}
if self.state_transition_times:
result["stateTransitionTimes"] = list(
self.state_transition_times
)
if self.recoverable_retries:
result["recoverableRetries"] = [
retry.as_dict() for retry in self.recoverable_retries
]
if (
self.jacobian_evaluation_count
or self.finite_difference_rhs_evaluation_count
or self.jacobian_assembly_seconds
):
result.update(
{
"jacobianEvaluationCount": self.jacobian_evaluation_count,
"jacobianFullBuildCount": self.jacobian_full_build_count,
"jacobianSecantReuseCount": self.jacobian_secant_reuse_count,
"jacobianAuditFailureCount": self.jacobian_audit_failure_count,
"finiteDifferenceRhsEvaluationCount": (
self.finite_difference_rhs_evaluation_count
),
"jacobianBaseRhsEvaluationCount": (
self.jacobian_base_rhs_evaluation_count
),
"jacobianJvAuditRhsEvaluationCount": (
self.jacobian_jv_audit_rhs_evaluation_count
),
"exactColumnBuildCount": self.exact_column_build_count,
"exactColumnFallbackCount": (
self.exact_column_fallback_count
),
"jacobianAssemblySeconds": self.jacobian_assembly_seconds,
}
)
return result
_JACOBIAN_DIAGNOSTIC_KEYS = (
"jacobianEvaluationCount",
"fullBuildCount",
"secantReuseCount",
"auditFailureCount",
"finiteDifferenceRhsEvaluationCount",
"baseRhsEvaluationCount",
"jvAuditEvaluationCount",
"exactColumnBuildCount",
"exactColumnFallbackCount",
"assemblySeconds",
)
def _jacobian_diagnostic_snapshot(
jac: JacobianCallable | None,
) -> dict[str, float]:
diagnostics = getattr(jac, "diagnostics", None)
if diagnostics is None:
return {key: 0.0 for key in _JACOBIAN_DIAGNOSTIC_KEYS}
values = diagnostics()
return {
key: float(values.get(key, 0.0))
for key in _JACOBIAN_DIAGNOSTIC_KEYS
}
def _positive_finite_step(value: object) -> float | None:
if value is None:
return None
try:
candidate = abs(float(value))
except (TypeError, ValueError, OverflowError):
return None
return candidate if candidate > 0.0 and math.isfinite(candidate) else None
def _smallest_positive_finite_step(*values: object) -> float:
"""Return a conservative step bound from configuration candidates."""
candidates = [
candidate
for value in values
if (candidate := _positive_finite_step(value)) is not None
]
if not candidates:
raise ValueError("No positive finite integration step is available.")
return min(candidates)
def _solver_attempted_step(
solver: object,
*,
segment_max_step: float,
remaining_interval: float,
) -> float:
"""Snapshot the real trial scale before calling ``solver.step()``.
SciPy exposes the proposed step as ``h_abs``. ``step_size`` is the prior
accepted step, so it is only a fallback for solvers without a valid
``h_abs``; it must not reduce an otherwise valid failed-trial estimate.
"""
configured_cap = _smallest_positive_finite_step(
segment_max_step,
remaining_interval,
)
for attribute in ("h_abs", "step_size"):
try:
candidate = _positive_finite_step(
getattr(solver, attribute, None)
)
except Exception:
# A third-party OdeSolver may implement these as fragile
# properties. The configured cap remains a safe fallback.
continue
if candidate is not None:
return min(candidate, configured_cap)
return configured_cap
def _recoverable_retry_steps(
attempted_step: float,
*,
last_accepted_time: float,
) -> tuple[float, float] | None:
"""Return strictly smaller max/first steps, or None at machine precision."""
next_step = _RECOVERABLE_RETRY_FACTOR * attempted_step
minimum_step = 64.0 * math.ulp(max(abs(last_accepted_time), 1.0))
if (
not math.isfinite(next_step)
or next_step <= minimum_step
or next_step >= attempted_step
):
return None
# This first step is intentionally one-shot. Keeping it equal to the new
# cap makes both controls strictly smaller than the failed trial scale.
return next_step, next_step
@dataclass(frozen=True)
class ODESolution:
t: list[float]
y: list[list[float]]
success: bool
message: str
status: IntegrationStatus = "completed"
error: Exception | None = None
solver_segments: tuple[SolverSegmentDiagnostics, ...] = ()
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:
time = float(time)
if times and time <= times[-1]:
return
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 _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,
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,
) -> list[float]:
"""Return sorted, unique breakpoints strictly inside the integration span."""
if breakpoints is None or len(breakpoints) == 0:
return []
if config.t_stop < config.t_start:
raise ValueError("Segmented integration requires t_stop to follow t_start.")
normalized: list[float] = []
for raw_breakpoint in breakpoints:
breakpoint = float(raw_breakpoint)
if not math.isfinite(breakpoint):
raise ValueError("Integration breakpoints must be finite numbers.")
if config.t_start < breakpoint < config.t_stop:
normalized.append(breakpoint)
normalized.sort()
return [
breakpoint
for index, breakpoint in enumerate(normalized)
if index == 0 or breakpoint != normalized[index - 1]
]
def _runge_kutta_4(
rhs: Callable[[float, list[float]], list[float]],
initial_state: list[float],
config: SolveIVPConfig,
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(
2,
int((config.t_stop - config.t_start) / max(config.max_step, 1e-6)) + 1,
)
step = (config.t_stop - config.t_start) / (point_count - 1)
t_eval = [config.t_start + index * step for index in range(point_count)]
state = list(initial_state)
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
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:
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))
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)
]
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:
status = "cancelled"
message = "Simulation was stopped before reaching the requested end time."
_append_solution_sample(times, states, current_time, state)
except Exception as exc:
status = "failed"
message = str(exc)
error = exc
_append_solution_sample(times, states, current_time, state)
return ODESolution(
t=times,
y=states,
success=status == "completed",
message=message,
status=status,
error=error,
)
def _runge_kutta_4_segmented(
rhs: Callable[[float, list[float]], list[float]],
initial_state: list[float],
config: SolveIVPConfig,
t_eval: list[float] | None,
breakpoints: Sequence[float],
cancel_check: CancellationCheck | None = None,
accepted_step_callback: AcceptedStepCallback | None = None,
state_transition_handler: StateTransitionHandler | None = None,
) -> ODESolution:
"""RK4 fallback that never evaluates a pre-breakpoint step at the breakpoint."""
if t_eval is None:
point_count = max(
2,
int((config.t_stop - config.t_start) / max(config.max_step, 1e-6)) + 1,
)
sample_step = (config.t_stop - config.t_start) / (point_count - 1)
sample_times = [
config.t_start + index * sample_step for index in range(point_count)
]
else:
sample_times = [float(time) for time in t_eval]
state = [float(value) for value in initial_state]
states = [[value] for value in state]
times = [float(config.t_start)]
current_time = float(config.t_start)
sample_index = 0
while (
sample_index < len(sample_times)
and sample_times[sample_index] <= config.t_start
):
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:
nonlocal last_reported_step
if accepted_step_callback is None:
return
if last_reported_step is not None and time <= last_reported_step:
return
accepted_step_callback(float(time))
last_reported_step = float(time)
def advance_to(
target_time: float, reported_terminal_time: float | None = None
) -> None:
nonlocal current_time, 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)
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))
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)
]
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
):
report_time = reported_terminal_time
report_step(report_time)
try:
segment_ends = [*breakpoints, float(config.t_stop)]
for segment_index, segment_end in enumerate(segment_ends):
is_breakpoint = segment_index < len(breakpoints)
integration_end = (
math.nextafter(segment_end, -math.inf) if is_breakpoint else segment_end
)
while (
sample_index < len(sample_times)
and sample_times[sample_index] <= integration_end
):
sample_time = float(sample_times[sample_index])
advance_to(sample_time)
_append_solution_sample(times, states, sample_time, state)
sample_index += 1
advance_to(
integration_end,
segment_end if is_breakpoint else None,
)
if is_breakpoint:
current_time = float(segment_end)
report_step(current_time)
while (
sample_index < len(sample_times)
and sample_times[sample_index] <= segment_end
):
sample_time = float(sample_times[sample_index])
_append_solution_sample(times, states, sample_time, state)
sample_index += 1
except IntegrationCancelled:
status = "cancelled"
message = "Simulation was stopped before reaching the requested end time."
_append_solution_sample(times, states, current_time, state)
except Exception as exc:
status = "failed"
message = str(exc)
error = exc
_append_solution_sample(times, states, current_time, state)
return ODESolution(
t=times,
y=states,
success=status == "completed",
message=message,
status=status,
error=error,
)
def _integrate_scipy_stepwise(
rhs: Callable[[float, list[float]], list[float]],
initial_state: list[float],
config: SolveIVPConfig,
t_eval: list[float] | None,
cancel_check: CancellationCheck,
accepted_step_callback: AcceptedStepCallback | None,
breakpoints: Sequence[float] = (),
state_transition_handler: StateTransitionHandler | None = None,
jac_sparsity=None,
jac: JacobianCallable | None = None,
) -> ODESolution:
"""Initial stepwise integration path for breakpoints and state resets.
Recoverable physical-domain failures from rejected integrator trial states
restore the last accepted state and rebuild the same solver with a smaller
maximum/first step. Structural, algebraic, and ordinary model errors still
fail immediately.
"""
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}")
implicit_jac = jac if config.method in {"BDF", "Radau"} else None
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 = [float(time) for time in (t_eval or [])]
sample_index = 0
while (
sample_index < len(sample_times)
and sample_times[sample_index] <= config.t_start
):
sample_index += 1
def cancellable_rhs(time, state):
if cancel_check():
raise IntegrationCancelled
normalized_state = [float(value) for value in state]
derivative = rhs(float(time), normalized_state)
observer = getattr(implicit_jac, "observe", None)
if observer is not None:
observer(float(time), normalized_state, derivative)
return derivative
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
solver_segments: list[SolverSegmentDiagnostics] = []
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:
return
if last_reported_step is not None and time <= last_reported_step:
return
accepted_step_callback(float(time))
last_reported_step = float(time)
segment_ends = [*breakpoints, float(config.t_stop)]
for segment_index, segment_end in enumerate(segment_ends):
if cancel_check():
status = "cancelled"
message = cancellation_message()
break
is_breakpoint = segment_index < len(breakpoints)
integration_end = (
math.nextafter(segment_end, -math.inf) if is_breakpoint else segment_end
)
has_integration_interval = integration_end > last_accepted_time
segment_start_time = last_accepted_time
segment_max_step = float(config.max_step)
recoverable_retry_count = 0
last_recoverable_error: RecoverableTrialStateError | None = None
retry_first_step: float | None = None
segment_nfev = 0
segment_njev = 0
segment_nlu = 0
segment_accepted_steps = 0
segment_solver_starts = 0
segment_state_transitions = 0
segment_state_transition_times: list[float] = []
segment_recoverable_retries = 0
segment_recoverable_retry_diagnostics: list[
RecoverableRetryDiagnostics
] = []
jacobian_work_start = _jacobian_diagnostic_snapshot(implicit_jac)
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,
"max_step": segment_max_step,
}
if config.method in {"BDF", "Radau"}:
if implicit_jac is not None:
solver_options["jac"] = implicit_jac
elif jac_sparsity is not None:
solver_options["jac_sparsity"] = jac_sparsity
requested_first_step = (
retry_first_step
if retry_first_step is not None
else config.first_step
)
if requested_first_step is not None:
solver_options["first_step"] = min(
requested_first_step,
integration_end - last_accepted_time,
)
try:
constructor_attempted_step = _smallest_positive_finite_step(
segment_max_step,
integration_end - last_accepted_time,
solver_options.get("first_step"),
)
start_segment = getattr(implicit_jac, "start_segment", None)
if start_segment is not None:
start_segment()
solver = solver_type(
cancellable_rhs,
last_accepted_time,
np.asarray(last_accepted_state, dtype=float),
integration_end,
**solver_options,
)
except IntegrationCancelled:
status = "cancelled"
message = cancellation_message()
break
except RecoverableTrialStateError as exc:
recoverable_retry_count += 1
segment_recoverable_retries += 1
last_recoverable_error = exc
retry_steps = (
_recoverable_retry_steps(
constructor_attempted_step,
last_accepted_time=last_accepted_time,
)
if recoverable_retry_count <= _MAX_RECOVERABLE_RETRIES
else None
)
segment_recoverable_retry_diagnostics.append(
RecoverableRetryDiagnostics(
phase="constructor",
attempted_step=constructor_attempted_step,
reason=str(exc),
next_max_step=(
retry_steps[0] if retry_steps is not None else None
),
next_first_step=(
retry_steps[1] if retry_steps is not None else None
),
)
)
if retry_steps is None:
status = "failed"
message = str(exc)
error = exc
break
segment_max_step, retry_first_step = retry_steps
continue
except Exception as exc:
status = "failed"
message = str(exc)
error = exc
break
segment_solver_starts += 1
restart_at_transition = False
restart_after_recoverable = False
while solver.status == "running":
if cancel_check():
status = "cancelled"
message = (
"Simulation was stopped before reaching the requested end time."
)
break
step_start_time = last_accepted_time
step_start_state = list(last_accepted_state)
try:
attempted_step = _solver_attempted_step(
solver,
segment_max_step=segment_max_step,
remaining_interval=(
integration_end - last_accepted_time
),
)
step_message = solver.step()
except IntegrationCancelled:
status = "cancelled"
message = (
"Simulation was stopped before reaching the requested end time."
)
break
except RecoverableTrialStateError as exc:
recoverable_retry_count += 1
segment_recoverable_retries += 1
last_recoverable_error = exc
retry_steps = (
_recoverable_retry_steps(
attempted_step,
last_accepted_time=last_accepted_time,
)
if recoverable_retry_count
<= _MAX_RECOVERABLE_RETRIES
else None
)
segment_recoverable_retry_diagnostics.append(
RecoverableRetryDiagnostics(
phase="step",
attempted_step=attempted_step,
reason=str(exc),
next_max_step=(
retry_steps[0]
if retry_steps is not None
else None
),
next_first_step=(
retry_steps[1]
if retry_steps is not None
else None
),
)
)
if retry_steps is None:
status = "failed"
message = str(exc)
error = exc
break
segment_max_step, retry_first_step = retry_steps
restart_after_recoverable = True
break
except Exception as exc:
status = "failed"
message = str(exc)
error = exc
break
integration_progressed = True
if solver.status == "failed":
if last_recoverable_error is not None:
recoverable_retry_count += 1
segment_recoverable_retries += 1
retry_steps = (
_recoverable_retry_steps(
attempted_step,
last_accepted_time=last_accepted_time,
)
if recoverable_retry_count
<= _MAX_RECOVERABLE_RETRIES
else None
)
failure_reason = str(
step_message or last_recoverable_error
)
segment_recoverable_retry_diagnostics.append(
RecoverableRetryDiagnostics(
phase="solver-status",
attempted_step=attempted_step,
reason=failure_reason,
next_max_step=(
retry_steps[0]
if retry_steps is not None
else None
),
next_first_step=(
retry_steps[1]
if retry_steps is not None
else None
),
)
)
if retry_steps is not None:
segment_max_step, retry_first_step = retry_steps
restart_after_recoverable = True
break
status = "failed"
message = str(step_message or "Integration step failed.")
break
# A returned running/finished status means this step was
# accepted. Any prior recoverable failure is now historical:
# it must not influence an event restart or an ordinary later
# solver failure. The reduced cap is local to the failed
# trial: after one accepted retry step, let this solver grow
# adaptively again and ensure a later event restart receives
# the configured maximum. The retry-specific first step is
# likewise strictly one-shot.
if retry_first_step is not None:
segment_max_step = float(config.max_step)
try:
solver.max_step = segment_max_step
except (AttributeError, TypeError, ValueError):
# Third-party OdeSolver-compatible test doubles may not
# expose a writable cap. SciPy's supported solvers do.
pass
last_recoverable_error = None
retry_first_step = None
recoverable_retry_count = 0
segment_accepted_steps += 1
step_end_time = float(solver.t)
step_end_state = [float(value) for value in solver.y]
crosses_sample = (
sample_index < len(sample_times)
and sample_times[sample_index] <= step_end_time
)
dense_output = (
solver.dense_output()
if crosses_sample 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:
segment_state_transitions += 1
segment_state_transition_times.append(
float(transition.time)
)
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
_project_nearby_pre_transition_sample(
times,
states,
transition,
config,
)
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:
while (
sample_index < len(sample_times)
and sample_times[sample_index] <= last_accepted_time
):
assert dense_output is not None
sample_time = float(sample_times[sample_index])
sample_state = [
float(value) for value in dense_output(sample_time)
]
_append_solution_sample(
times,
states,
sample_time,
sample_state,
)
sample_index += 1
else:
_append_solution_sample(
times,
states,
reported_time,
last_accepted_state,
)
report_step(reported_time)
segment_nfev += int(getattr(solver, "nfev", 0))
segment_njev += int(getattr(solver, "njev", 0))
segment_nlu += int(getattr(solver, "nlu", 0))
if status != "completed":
break
if restart_after_recoverable:
continue
if not restart_at_transition:
break
jacobian_work_end = _jacobian_diagnostic_snapshot(implicit_jac)
jacobian_work = {
key: jacobian_work_end[key] - jacobian_work_start[key]
for key in _JACOBIAN_DIAGNOSTIC_KEYS
}
solver_segments.append(
SolverSegmentDiagnostics(
start_time=float(segment_start_time),
requested_stop_time=float(segment_end),
simulated_until=float(
segment_end if status == "completed" else last_accepted_time
),
nfev=segment_nfev,
njev=segment_njev,
nlu=segment_nlu,
accepted_step_count=segment_accepted_steps,
solver_start_count=segment_solver_starts,
state_transition_count=segment_state_transitions,
state_transition_times=tuple(
segment_state_transition_times
),
recoverable_retry_count=segment_recoverable_retries,
recoverable_retries=tuple(
segment_recoverable_retry_diagnostics
),
jacobian_evaluation_count=int(
jacobian_work["jacobianEvaluationCount"]
),
jacobian_full_build_count=int(
jacobian_work["fullBuildCount"]
),
jacobian_secant_reuse_count=int(
jacobian_work["secantReuseCount"]
),
jacobian_audit_failure_count=int(
jacobian_work["auditFailureCount"]
),
finite_difference_rhs_evaluation_count=int(
jacobian_work["finiteDifferenceRhsEvaluationCount"]
),
jacobian_base_rhs_evaluation_count=int(
jacobian_work["baseRhsEvaluationCount"]
),
jacobian_jv_audit_rhs_evaluation_count=int(
jacobian_work["jvAuditEvaluationCount"]
),
exact_column_build_count=int(
jacobian_work["exactColumnBuildCount"]
),
exact_column_fallback_count=int(
jacobian_work["exactColumnFallbackCount"]
),
jacobian_assembly_seconds=jacobian_work["assemblySeconds"],
)
)
if status != "completed":
break
if is_breakpoint:
# The old equation is integrated only to the representable point just
# left of the event. The continuous state is then lifted to the exact
# event time, where the freshly constructed next solver sees the new
# equation immediately.
last_accepted_time = float(segment_end)
if sample_times:
while (
sample_index < len(sample_times)
and sample_times[sample_index] <= segment_end
):
sample_time = float(sample_times[sample_index])
_append_solution_sample(
times,
states,
sample_time,
last_accepted_state,
)
sample_index += 1
elif not has_integration_interval:
_append_solution_sample(
times,
states,
last_accepted_time,
last_accepted_state,
)
report_step(last_accepted_time)
if status != "completed":
_append_solution_sample(
times,
states,
last_accepted_time,
last_accepted_state,
)
return ODESolution(
t=times,
y=states,
success=status == "completed",
message=message,
status=status,
error=error,
solver_segments=tuple(solver_segments),
)
@profile_phase("simulation.integration")
def integrate_ode(
rhs: Callable[[float, list[float]], list[float]],
initial_state: list[float],
config: SolveIVPConfig,
t_eval: list[float] | None = None,
cancel_check: CancellationCheck | None = None,
accepted_step_callback: AcceptedStepCallback | None = None,
breakpoints: Sequence[float] | None = None,
state_transition_handler: StateTransitionHandler | None = None,
jac_sparsity=None,
jac: JacobianCallable | None = None,
recoverable_trial_retries: bool = False,
):
"""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.
``recoverable_trial_retries`` opts an eventless/cancellation-free caller
into the stepwise path so a ``RecoverableTrialStateError`` can rebuild the
solver from its last accepted state. It defaults to false to preserve the
direct ``solve_ivp`` path for ordinary callers.
"""
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],
success=True,
message="Skipped integration because t_start equals t_stop.",
)
normalized_breakpoints = _normalize_breakpoints(config, breakpoints)
try:
from scipy.integrate import solve_ivp
except ImportError:
if normalized_breakpoints:
return _runge_kutta_4_segmented(
rhs,
initial_state,
config,
t_eval,
normalized_breakpoints,
cancel_check,
accepted_step_callback,
state_transition_handler,
)
return _runge_kutta_4(
rhs,
initial_state,
config,
t_eval,
cancel_check,
accepted_step_callback,
state_transition_handler,
)
if (
cancel_check is not None
or normalized_breakpoints
or state_transition_handler is not None
or recoverable_trial_retries
):
return _integrate_scipy_stepwise(
rhs,
initial_state,
config,
t_eval,
cancel_check or (lambda: False),
accepted_step_callback,
normalized_breakpoints,
state_transition_handler,
jac_sparsity,
jac,
)
implicit_jac = jac if config.method in {"BDF", "Radau"} else None
solve_rhs = rhs
if implicit_jac is not None:
observer = getattr(implicit_jac, "observe", None)
if observer is not None:
def observed_rhs(time, state):
derivative = rhs(time, state)
observer(float(time), state, derivative)
return derivative
solve_rhs = observed_rhs
start_segment = getattr(implicit_jac, "start_segment", None)
if start_segment is not None:
start_segment()
solve_options = {
"fun": solve_rhs,
"t_span": (config.t_start, config.t_stop),
"y0": initial_state,
"method": config.method,
"rtol": config.rtol,
"atol": config.atol,
"max_step": config.max_step,
"t_eval": t_eval,
}
if config.first_step is not None:
solve_options["first_step"] = config.first_step
if implicit_jac is not None:
solve_options["jac"] = implicit_jac
elif jac_sparsity is not None and config.method in {"BDF", "Radau"}:
solve_options["jac_sparsity"] = jac_sparsity
return solve_ivp(**solve_options)