优化仿真求解性能并修复流量闭合问题(初版)

This commit is contained in:
ljz committed 2026-08-16 17:46:05 +08:00
1 parent 57b459bc72
commit 5332a788f3
55 files changed
+8973 -549

No files matched your search

+54 -20
View File
@@ -2,8 +2,10 @@ from __future__ import annotations
from dataclasses import dataclass
from math import isfinite
from typing import Protocol
from typing import Callable, Protocol
from app.simulation.core.base import Component
from app.simulation.core.ports import PortState
from app.simulation.performance import profile_phase
from app.simulation.systems.network import Endpoint, SimulationNetwork
@@ -38,31 +40,56 @@ class SignalSolveDiagnostics:
return {"propagated": self.propagated}
@dataclass(frozen=True)
class _SignalOutputBinding:
component: Component
evaluate: Callable[[float], dict[str, float]]
@dataclass(frozen=True)
class _SignalConnectionBinding:
source: PortState
target: PortState
class SignalResolver:
"""Propagate scalar signal connections from output ports to input ports."""
def __init__(self, network: SimulationNetwork) -> None:
self.network = network
self._connections = [
connection for connection in network.connections if connection.kind == "signal"
]
self._output_bindings = tuple(
_SignalOutputBinding(component=component, evaluate=evaluate)
for component in network.components.values()
if (evaluate := getattr(component, "signal_output_values", None)) is not None
)
self._event_sources = tuple(
(component.name, source_event_times)
for component in network.components.values()
if (
source_event_times := getattr(
component,
"signal_event_times",
None,
)
)
is not None
)
self._connections = tuple(
self._connection_binding(connection.endpoints)
for connection in network.connections
if connection.kind == "signal"
)
self.last_diagnostics: SignalSolveDiagnostics | None = None
@profile_phase("simulation.signal", minimum_mode="audit")
def solve(self, time: float) -> SignalSolveDiagnostics:
for component in self.network.components.values():
signal_output_values = getattr(component, "signal_output_values", None)
if signal_output_values is None:
continue
for port_name, value in signal_output_values(time).items():
component.get_port(port_name).signal = float(value)
for binding in self._output_bindings:
for port_name, value in binding.evaluate(time).items():
binding.component.get_port(port_name).signal = float(value)
propagated = 0
for connection in self._connections:
source, target = self._source_target(connection.endpoints)
source_port = self.network.components[source.component].get_port(source.port)
target_port = self.network.components[target.component].get_port(target.port)
target_port.signal = source_port.signal
for binding in self._connections:
binding.target.signal = binding.source.signal
propagated += 1
diagnostics = SignalSolveDiagnostics(propagated=propagated)
@@ -86,15 +113,12 @@ class SignalResolver:
return ()
events: set[float] = set()
for component in self.network.components.values():
source_event_times = getattr(component, "signal_event_times", None)
if source_event_times is None:
continue
for component_name, source_event_times in self._event_sources:
for raw_time in source_event_times(start, stop):
event_time = float(raw_time)
if not isfinite(event_time):
raise ValueError(
f"Signal event time from component '{component.name}' must be finite."
f"Signal event time from component '{component_name}' must be finite."
)
if start < event_time < stop:
events.add(event_time)
@@ -109,3 +133,13 @@ class SignalResolver:
if second_port.definition is not None and second_port.definition.nominal_role == "output":
return second, first
raise ValueError("Signal connection must contain one output endpoint.")
def _connection_binding(
self,
endpoints: tuple[Endpoint, Endpoint],
) -> _SignalConnectionBinding:
source, target = self._source_target(endpoints)
return _SignalConnectionBinding(
source=self.network.components[source.component].get_port(source.port),
target=self.network.components[target.component].get_port(target.port),
)