初版:实现 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

+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 = {