568 lines
19 KiB
Python
568 lines
19 KiB
Python
from __future__ import annotations
|
|
|
|
from collections.abc import Callable, Sequence
|
|
from copy import copy
|
|
from dataclasses import dataclass, replace
|
|
|
|
from app.simulation.core.errors import RecoverableTrialStateError
|
|
from app.simulation.core.ports import PortState
|
|
|
|
|
|
_STREAM_CACHE_ATTRIBUTE_NAMES = frozenset(
|
|
{
|
|
"_connected_h",
|
|
"temperature_reference_h",
|
|
}
|
|
)
|
|
|
|
|
|
def _is_stream_cache_attribute(name: str) -> bool:
|
|
"""Return whether an attribute belongs to the stream/temperature replay state.
|
|
|
|
Catalog components currently use ``_connected_h`` and
|
|
``temperature_reference_h``. The name-based extension keeps conservative
|
|
third-party caches recoverable without copying an entire component graph.
|
|
Components with opaque cache names can provide the explicit hooks documented
|
|
by :class:`ThermofluidTransactionPlan`.
|
|
"""
|
|
|
|
lowered = name.lower()
|
|
return (
|
|
name in _STREAM_CACHE_ATTRIBUTE_NAMES
|
|
or lowered.startswith("_stream_")
|
|
or "connected_h" in lowered
|
|
or "connected_enthalpy" in lowered
|
|
or "temperature_reference" in lowered
|
|
)
|
|
|
|
|
|
def _copy_cache_value(value: object) -> object:
|
|
"""Shallow-copy a stream cache without traversing the component graph."""
|
|
|
|
if isinstance(value, (dict, list, set, bytearray)):
|
|
return copy(value)
|
|
return value
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class ThermofluidWorstPort:
|
|
component: str
|
|
port: str
|
|
value: float
|
|
signed_delta: float
|
|
|
|
def as_dict(self) -> dict[str, object]:
|
|
return {
|
|
"component": self.component,
|
|
"port": self.port,
|
|
"value": self.value,
|
|
"signedDelta": self.signed_delta,
|
|
}
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class ThermofluidIterationDelta:
|
|
iteration: int
|
|
max_delta: float
|
|
scale: float
|
|
tolerance: float
|
|
worst_port: ThermofluidWorstPort | None
|
|
|
|
def as_dict(self) -> dict[str, object]:
|
|
return {
|
|
"iteration": self.iteration,
|
|
"maxDelta": self.max_delta,
|
|
"scale": self.scale,
|
|
"tolerance": self.tolerance,
|
|
"worstPort": (
|
|
self.worst_port.as_dict()
|
|
if self.worst_port is not None
|
|
else None
|
|
),
|
|
}
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class ThermofluidClosureSuccess:
|
|
rhs_time: float
|
|
iterations: int
|
|
max_delta: float
|
|
scale: float
|
|
tolerance: float
|
|
worst_port: ThermofluidWorstPort | None
|
|
|
|
@classmethod
|
|
def from_iteration(
|
|
cls,
|
|
rhs_time: float,
|
|
delta: ThermofluidIterationDelta,
|
|
) -> ThermofluidClosureSuccess:
|
|
return cls(
|
|
rhs_time=float(rhs_time),
|
|
iterations=delta.iteration,
|
|
max_delta=delta.max_delta,
|
|
scale=delta.scale,
|
|
tolerance=delta.tolerance,
|
|
worst_port=delta.worst_port,
|
|
)
|
|
|
|
def as_dict(self) -> dict[str, object]:
|
|
return {
|
|
"rhsTime": self.rhs_time,
|
|
"iterations": self.iterations,
|
|
"maxDelta": self.max_delta,
|
|
"scale": self.scale,
|
|
"tolerance": self.tolerance,
|
|
"worstPort": (
|
|
self.worst_port.as_dict()
|
|
if self.worst_port is not None
|
|
else None
|
|
),
|
|
}
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class ThermofluidClosureFailure:
|
|
failed_rhs_time: float
|
|
iterations: int
|
|
delta_tail: tuple[ThermofluidIterationDelta, ...]
|
|
max_delta: float
|
|
scale: float
|
|
tolerance: float
|
|
worst_port: ThermofluidWorstPort | None
|
|
failure_count: int = 0
|
|
|
|
@classmethod
|
|
def from_iterations(
|
|
cls,
|
|
failed_rhs_time: float,
|
|
deltas: Sequence[ThermofluidIterationDelta],
|
|
*,
|
|
tail_limit: int = 8,
|
|
) -> ThermofluidClosureFailure:
|
|
if not deltas:
|
|
raise ValueError("A thermofluid failure requires iteration diagnostics.")
|
|
final = deltas[-1]
|
|
return cls(
|
|
failed_rhs_time=float(failed_rhs_time),
|
|
iterations=final.iteration,
|
|
delta_tail=tuple(deltas[-tail_limit:]),
|
|
max_delta=final.max_delta,
|
|
scale=final.scale,
|
|
tolerance=final.tolerance,
|
|
worst_port=final.worst_port,
|
|
)
|
|
|
|
def as_dict(self) -> dict[str, object]:
|
|
return {
|
|
"failedRhsTime": self.failed_rhs_time,
|
|
"iterations": self.iterations,
|
|
"deltaTail": [item.as_dict() for item in self.delta_tail],
|
|
"maxDelta": self.max_delta,
|
|
"scale": self.scale,
|
|
"tolerance": self.tolerance,
|
|
"worstPort": (
|
|
self.worst_port.as_dict()
|
|
if self.worst_port is not None
|
|
else None
|
|
),
|
|
"failureCount": self.failure_count,
|
|
}
|
|
|
|
|
|
class ThermofluidClosureError(RecoverableTrialStateError):
|
|
"""Recoverable exhaustion of the stream/pressure-flow fixed point.
|
|
|
|
Stream propagation failures and algebraic-solver failures intentionally
|
|
retain their original exception types: rollback is still applied, but a
|
|
smaller ODE step is not known to repair those structural/numerical errors.
|
|
"""
|
|
|
|
def __init__(self, diagnostics: ThermofluidClosureFailure) -> None:
|
|
super().__init__(
|
|
"Stream enthalpy and pressure-flow coupling did not converge "
|
|
f"after {diagnostics.iterations} iterations at "
|
|
f"t={diagnostics.failed_rhs_time:.17g}."
|
|
)
|
|
self.diagnostics = diagnostics
|
|
|
|
|
|
class ThermofluidClosureDiagnostics:
|
|
"""Run-level RHS outcomes; maintenance/postprocessing calls do not write it."""
|
|
|
|
def __init__(self) -> None:
|
|
self.failure_count = 0
|
|
self.last_failure: ThermofluidClosureFailure | None = None
|
|
self.last_success: ThermofluidClosureSuccess | None = None
|
|
|
|
def record_success(self, success: ThermofluidClosureSuccess) -> None:
|
|
self.last_success = success
|
|
|
|
def record_failure(
|
|
self,
|
|
failure: ThermofluidClosureFailure,
|
|
) -> ThermofluidClosureFailure:
|
|
self.failure_count += 1
|
|
recorded = replace(failure, failure_count=self.failure_count)
|
|
self.last_failure = recorded
|
|
return recorded
|
|
|
|
def as_dict(self) -> dict[str, object]:
|
|
return {
|
|
"failureCount": self.failure_count,
|
|
"lastFailure": (
|
|
self.last_failure.as_dict()
|
|
if self.last_failure is not None
|
|
else None
|
|
),
|
|
"lastSuccess": (
|
|
self.last_success.as_dict()
|
|
if self.last_success is not None
|
|
else None
|
|
),
|
|
}
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class _PortValueBinding:
|
|
component_name: str
|
|
port_name: str
|
|
state: PortState
|
|
variable: str
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class _PortFieldPlan:
|
|
variable: str
|
|
states: tuple[PortState, ...]
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class _FlowBinding:
|
|
component_name: str
|
|
port_name: str
|
|
state: PortState
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class _ComponentCacheBinding:
|
|
component: object
|
|
attribute_names: tuple[str, ...]
|
|
attribute_name_set: frozenset[str]
|
|
snapshot_hook: Callable[[], object] | None
|
|
restore_hook: Callable[[object], None] | None
|
|
|
|
|
|
@dataclass
|
|
class ThermofluidTransactionSnapshot:
|
|
plan: ThermofluidTransactionPlan
|
|
port_values: tuple[list[float], ...]
|
|
component_cache_values: tuple[list[object], ...]
|
|
custom_cache_values: list[object | None]
|
|
diagnostic_values: list[object]
|
|
|
|
def restore(self) -> None:
|
|
plan = self.plan
|
|
plan._restore_port_values(self.port_values)
|
|
|
|
for binding, values, custom_value in zip(
|
|
plan.component_cache_bindings,
|
|
self.component_cache_values,
|
|
self.custom_cache_values,
|
|
):
|
|
component = binding.component
|
|
for name in tuple(getattr(component, "__dict__", {})):
|
|
if (
|
|
name.startswith("_causal_")
|
|
or _is_stream_cache_attribute(name)
|
|
) and name not in binding.attribute_name_set:
|
|
delattr(component, name)
|
|
for name, value in zip(binding.attribute_names, values):
|
|
setattr(component, name, _copy_cache_value(value))
|
|
if binding.restore_hook is not None:
|
|
binding.restore_hook(custom_value)
|
|
|
|
for owner, value in zip(
|
|
plan.diagnostic_owners,
|
|
self.diagnostic_values,
|
|
):
|
|
owner.last_diagnostics = value
|
|
|
|
|
|
class ThermofluidTransactionPlan:
|
|
"""Compiled, lightweight rollback boundary for one Generic RHS closure.
|
|
|
|
It snapshots active physical-port values, catalog stream-temperature caches,
|
|
component ``_causal_*`` seed fields, and resolver/solver last diagnostics.
|
|
A custom stream-aware component with an opaque mutable cache can implement
|
|
both ``snapshot_thermofluid_closure_cache()`` and
|
|
``restore_thermofluid_closure_cache(snapshot)``; these hooks are invoked in
|
|
addition to the standard name-based cache capture.
|
|
"""
|
|
|
|
def __init__(
|
|
self,
|
|
*,
|
|
port_value_bindings: tuple[_PortValueBinding, ...],
|
|
port_field_plans: tuple[_PortFieldPlan, ...],
|
|
flow_bindings: tuple[_FlowBinding, ...],
|
|
component_cache_bindings: tuple[_ComponentCacheBinding, ...],
|
|
component_count: int,
|
|
diagnostic_owners: tuple[object, ...],
|
|
) -> None:
|
|
self.port_value_bindings = port_value_bindings
|
|
self.port_field_plans = port_field_plans
|
|
self.flow_bindings = flow_bindings
|
|
self.component_cache_bindings = component_cache_bindings
|
|
self.component_count = component_count
|
|
self.diagnostic_owners = diagnostic_owners
|
|
self._snapshot = ThermofluidTransactionSnapshot(
|
|
plan=self,
|
|
port_values=tuple(
|
|
[0.0] * len(field.states)
|
|
for field in port_field_plans
|
|
),
|
|
component_cache_values=tuple(
|
|
[None] * len(binding.attribute_names)
|
|
for binding in component_cache_bindings
|
|
),
|
|
custom_cache_values=[None] * len(component_cache_bindings),
|
|
diagnostic_values=[None] * len(diagnostic_owners),
|
|
)
|
|
|
|
@classmethod
|
|
def compile(
|
|
cls,
|
|
network: object,
|
|
*,
|
|
diagnostic_owners: Sequence[object] = (),
|
|
) -> ThermofluidTransactionPlan:
|
|
components = tuple(getattr(network, "components").values())
|
|
port_value_bindings: list[_PortValueBinding] = []
|
|
port_states_by_variable: dict[str, list[PortState]] = {}
|
|
flow_bindings: list[_FlowBinding] = []
|
|
component_cache_bindings: list[_ComponentCacheBinding] = []
|
|
for component in components:
|
|
active_definitions = tuple(
|
|
definition
|
|
for definition in component.active_port_definitions
|
|
if definition.kind == "physical"
|
|
)
|
|
for definition in active_definitions:
|
|
state = component.get_port(definition.name)
|
|
flow_bindings.append(
|
|
_FlowBinding(component.name, definition.name, state)
|
|
)
|
|
for variable in definition.variables:
|
|
port_states_by_variable.setdefault(variable.name, []).append(state)
|
|
port_value_bindings.append(
|
|
_PortValueBinding(
|
|
component.name,
|
|
definition.name,
|
|
state,
|
|
variable.name,
|
|
)
|
|
)
|
|
|
|
attribute_names = tuple(
|
|
name
|
|
for name in getattr(component, "__dict__", {})
|
|
if name.startswith("_causal_")
|
|
or _is_stream_cache_attribute(name)
|
|
)
|
|
snapshot_hook = getattr(
|
|
component,
|
|
"snapshot_thermofluid_closure_cache",
|
|
None,
|
|
)
|
|
restore_hook = getattr(
|
|
component,
|
|
"restore_thermofluid_closure_cache",
|
|
None,
|
|
)
|
|
hooks_are_available = callable(snapshot_hook) and callable(restore_hook)
|
|
if attribute_names or hooks_are_available:
|
|
component_cache_bindings.append(
|
|
_ComponentCacheBinding(
|
|
component=component,
|
|
attribute_names=attribute_names,
|
|
attribute_name_set=frozenset(attribute_names),
|
|
snapshot_hook=(snapshot_hook if hooks_are_available else None),
|
|
restore_hook=(restore_hook if hooks_are_available else None),
|
|
)
|
|
)
|
|
|
|
owners = tuple(
|
|
dict.fromkeys(
|
|
owner
|
|
for owner in diagnostic_owners
|
|
if hasattr(owner, "last_diagnostics")
|
|
)
|
|
)
|
|
return cls(
|
|
port_value_bindings=tuple(port_value_bindings),
|
|
port_field_plans=tuple(
|
|
_PortFieldPlan(variable, tuple(states))
|
|
for variable, states in port_states_by_variable.items()
|
|
),
|
|
flow_bindings=tuple(flow_bindings),
|
|
component_cache_bindings=tuple(component_cache_bindings),
|
|
component_count=len(components),
|
|
diagnostic_owners=owners,
|
|
)
|
|
|
|
def capture(self) -> ThermofluidTransactionSnapshot:
|
|
# GenericFluidSystem executes one RHS serially. Reuse one compiled
|
|
# workspace rather than allocating a snapshot object and several outer
|
|
# tuples at every successful trial point.
|
|
snapshot = self._snapshot
|
|
self._capture_port_values(snapshot.port_values)
|
|
for binding, values in zip(
|
|
self.component_cache_bindings,
|
|
snapshot.component_cache_values,
|
|
):
|
|
for position, name in enumerate(binding.attribute_names):
|
|
values[position] = _copy_cache_value(
|
|
getattr(binding.component, name)
|
|
)
|
|
for position, binding in enumerate(self.component_cache_bindings):
|
|
snapshot.custom_cache_values[position] = (
|
|
binding.snapshot_hook()
|
|
if binding.snapshot_hook is not None
|
|
else None
|
|
)
|
|
for position, owner in enumerate(self.diagnostic_owners):
|
|
snapshot.diagnostic_values[position] = owner.last_diagnostics
|
|
return snapshot
|
|
|
|
def _capture_port_values(
|
|
self,
|
|
workspaces: tuple[list[float], ...],
|
|
) -> None:
|
|
for field, values in zip(self.port_field_plans, workspaces):
|
|
variable = field.variable
|
|
states = field.states
|
|
if variable == "p":
|
|
for position, state in enumerate(states):
|
|
values[position] = state.p
|
|
elif variable == "m_flow":
|
|
for position, state in enumerate(states):
|
|
values[position] = state.m_flow
|
|
elif variable == "h_outflow":
|
|
for position, state in enumerate(states):
|
|
values[position] = state.h_outflow
|
|
elif variable == "volume":
|
|
for position, state in enumerate(states):
|
|
values[position] = state.volume
|
|
elif variable == "volume_flow":
|
|
for position, state in enumerate(states):
|
|
values[position] = state.volume_flow
|
|
elif variable == "x":
|
|
for position, state in enumerate(states):
|
|
values[position] = state.x
|
|
elif variable == "v":
|
|
for position, state in enumerate(states):
|
|
values[position] = state.v
|
|
elif variable == "f":
|
|
for position, state in enumerate(states):
|
|
values[position] = state.f
|
|
else:
|
|
for position, state in enumerate(states):
|
|
values[position] = getattr(state, variable)
|
|
|
|
def _restore_port_values(
|
|
self,
|
|
workspaces: tuple[list[float], ...],
|
|
) -> None:
|
|
for field, values in zip(self.port_field_plans, workspaces):
|
|
variable = field.variable
|
|
states = field.states
|
|
if variable == "p":
|
|
for state, value in zip(states, values):
|
|
state.p = value
|
|
elif variable == "m_flow":
|
|
for state, value in zip(states, values):
|
|
state.m_flow = value
|
|
elif variable == "h_outflow":
|
|
for state, value in zip(states, values):
|
|
state.h_outflow = value
|
|
elif variable == "volume":
|
|
for state, value in zip(states, values):
|
|
state.volume = value
|
|
elif variable == "volume_flow":
|
|
for state, value in zip(states, values):
|
|
state.volume_flow = value
|
|
elif variable == "x":
|
|
for state, value in zip(states, values):
|
|
state.x = value
|
|
elif variable == "v":
|
|
for state, value in zip(states, values):
|
|
state.v = value
|
|
elif variable == "f":
|
|
for state, value in zip(states, values):
|
|
state.f = value
|
|
else:
|
|
for state, value in zip(states, values):
|
|
setattr(state, variable, value)
|
|
|
|
def flow_values(self) -> tuple[float, ...]:
|
|
return tuple(float(binding.state.m_flow) for binding in self.flow_bindings)
|
|
|
|
def measure_flow_delta(
|
|
self,
|
|
previous: Sequence[float],
|
|
*,
|
|
iteration: int,
|
|
relative_tolerance: float,
|
|
) -> ThermofluidIterationDelta:
|
|
current = self.flow_values()
|
|
scale = max(
|
|
(abs(value) for value in (*previous, *current)),
|
|
default=1.0,
|
|
)
|
|
scale = max(scale, 1.0)
|
|
worst_index = -1
|
|
worst_signed_delta = 0.0
|
|
max_delta = 0.0
|
|
for index, (old, new) in enumerate(zip(previous, current)):
|
|
signed_delta = new - old
|
|
magnitude = abs(signed_delta)
|
|
if magnitude > max_delta:
|
|
worst_index = index
|
|
worst_signed_delta = signed_delta
|
|
max_delta = magnitude
|
|
worst_port = None
|
|
if worst_index >= 0:
|
|
binding = self.flow_bindings[worst_index]
|
|
worst_port = ThermofluidWorstPort(
|
|
component=binding.component_name,
|
|
port=binding.port_name,
|
|
value=current[worst_index],
|
|
signed_delta=worst_signed_delta,
|
|
)
|
|
return ThermofluidIterationDelta(
|
|
iteration=int(iteration),
|
|
max_delta=max_delta,
|
|
scale=scale,
|
|
tolerance=float(relative_tolerance) * scale,
|
|
worst_port=worst_port,
|
|
)
|
|
|
|
def diagnostics(self) -> dict[str, int]:
|
|
stream_cache_slot_count = sum(
|
|
len(binding.attribute_names)
|
|
for binding in self.component_cache_bindings
|
|
)
|
|
return {
|
|
"physicalPortValueSlotCount": len(self.port_value_bindings),
|
|
"physicalFlowPortCount": len(self.flow_bindings),
|
|
"componentCount": self.component_count,
|
|
"cacheBindingCount": len(self.component_cache_bindings),
|
|
"streamAndCausalCacheSlotCount": stream_cache_slot_count,
|
|
"customCacheHookCount": sum(
|
|
binding.snapshot_hook is not None
|
|
for binding in self.component_cache_bindings
|
|
),
|
|
"diagnosticOwnerCount": len(self.diagnostic_owners),
|
|
}
|