优化仿真求解性能并修复流量闭合问题(初版)
This commit is contained in:
1 parent
57b459bc72
commit
5332a788f3
55 files changed
+8973
-549
No files matched your search
@@ -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),
|
||||
)
|
||||
Reference in new issue
Block a user