258 lines
12 KiB
Python
258 lines
12 KiB
Python
"""Dependency ordering and local algebraic blocks for generated C expressions.
|
|
|
|
This module only arranges reviewed C computations. It never evaluates a model
|
|
numerically and has no dependency on the retired Python numerical backend.
|
|
"""
|
|
from __future__ import annotations
|
|
|
|
from dataclasses import dataclass
|
|
import heapq
|
|
import re
|
|
|
|
from .compiler import NativeCapabilityError
|
|
|
|
|
|
# Only compiler-owned array expressions are inspected, never arbitrary user C.
|
|
_REFERENCE = re.compile(r"\b(?:[phqw]|fb)\[\d+\]|\bg\[\d+\]\.[A-Za-z_]\w*")
|
|
|
|
|
|
def references(expression: str) -> frozenset[str]:
|
|
return frozenset(_REFERENCE.findall(expression))
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class Computation:
|
|
key: str
|
|
outputs: tuple[str, ...]
|
|
inputs: frozenset[str]
|
|
code: tuple[str, ...]
|
|
kind: str = "flow"
|
|
residual: str | None = None
|
|
|
|
@classmethod
|
|
def assignment(cls, key, target, expression, kind="flow"):
|
|
return cls(key, (target,), references(expression), (f'{target}={expression};',), kind)
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class Block:
|
|
members: tuple[int, ...]
|
|
cyclic: bool
|
|
|
|
|
|
class EvaluationSchedule:
|
|
def __init__(self, computations, known, labels=None):
|
|
self.computations = tuple(computations)
|
|
self.labels = labels or {}
|
|
self.producers = {}
|
|
for i, op in enumerate(self.computations):
|
|
for output in op.outputs:
|
|
if output in self.producers:
|
|
raise NativeCapabilityError(f'Multiple native producers for {self.label(output)}')
|
|
self.producers[output] = i
|
|
# An iterative initial guess is not a known source if an equation owns it.
|
|
used = set().union(*(op.inputs for op in self.computations)) if self.computations else set()
|
|
self.known = {key: value for key, value in known.items() if key not in self.producers and key in used}
|
|
self.dependencies = []
|
|
for op in self.computations:
|
|
missing = op.inputs - self.producers.keys() - self.known.keys()
|
|
if missing:
|
|
raise NativeCapabilityError(f'{op.key}: missing native input sources: {sorted(map(self.label, missing))}')
|
|
self.dependencies.append({self.producers[key] for key in op.inputs if key in self.producers})
|
|
self.blocks = self._blocks()
|
|
|
|
def label(self, key):
|
|
return self.labels.get(key, key)
|
|
|
|
def _blocks(self):
|
|
"""Iterative SCC discovery followed by deterministic topological order."""
|
|
count = len(self.computations)
|
|
consumers = [set() for _ in range(count)]
|
|
for target, sources in enumerate(self.dependencies):
|
|
for source in sources:
|
|
consumers[source].add(target)
|
|
visited, finish = set(), []
|
|
for start in range(count):
|
|
if start in visited:
|
|
continue
|
|
visited.add(start)
|
|
stack = [(start, iter(sorted(consumers[start])))]
|
|
while stack:
|
|
node, edges = stack[-1]
|
|
child = next(edges, None)
|
|
if child is None:
|
|
finish.append(node)
|
|
stack.pop()
|
|
elif child not in visited:
|
|
visited.add(child)
|
|
stack.append((child, iter(sorted(consumers[child]))))
|
|
groups, owner = [], {}
|
|
for start in reversed(finish):
|
|
if start in owner:
|
|
continue
|
|
index = len(groups)
|
|
owner[start] = index
|
|
members, stack = [], [start]
|
|
while stack:
|
|
node = stack.pop()
|
|
members.append(node)
|
|
for child in sorted(self.dependencies[node]):
|
|
if child not in owner:
|
|
owner[child] = index
|
|
stack.append(child)
|
|
groups.append(tuple(sorted(members)))
|
|
incoming = [set() for _ in groups]
|
|
outgoing = [set() for _ in groups]
|
|
for target, sources in enumerate(self.dependencies):
|
|
for source in sources:
|
|
a, b = owner[source], owner[target]
|
|
if a != b:
|
|
incoming[b].add(a)
|
|
outgoing[a].add(b)
|
|
ready = [(min(groups[i]), i) for i, inputs in enumerate(incoming) if not inputs]
|
|
heapq.heapify(ready)
|
|
blocks = []
|
|
while ready:
|
|
_, index = heapq.heappop(ready)
|
|
members = groups[index]
|
|
cyclic = len(members) > 1 or members[0] in self.dependencies[members[0]]
|
|
blocks.append(Block(members, cyclic))
|
|
for child in sorted(outgoing[index]):
|
|
incoming[child].remove(index)
|
|
if not incoming[child]:
|
|
heapq.heappush(ready, (min(groups[child]), child))
|
|
return tuple(blocks)
|
|
|
|
def ordered_subset(self, members):
|
|
"""Order a trial's computations while pressure/stream guesses are fixed."""
|
|
pending = set(members)
|
|
result = []
|
|
while pending:
|
|
ready = sorted(i for i in pending if not self.dependencies[i] & pending)
|
|
if not ready:
|
|
raise NativeCapabilityError('Unsupported cycle within native flow expressions')
|
|
result.extend(ready)
|
|
pending.difference_update(ready)
|
|
return result
|
|
|
|
def ancestors(self, inputs, allowed):
|
|
pending = [self.producers[key] for key in inputs if key in self.producers]
|
|
found = set()
|
|
while pending:
|
|
index = pending.pop()
|
|
if index in found or index not in allowed:
|
|
continue
|
|
found.add(index)
|
|
pending.extend(self.dependencies[index])
|
|
return self.ordered_subset(found)
|
|
|
|
def report(self):
|
|
result = []
|
|
sources = {key: {str(origin)} for key, origin in self.known.items()}
|
|
for block in self.blocks:
|
|
outputs = {key for i in block.members for key in self.computations[i].outputs}
|
|
inputs = {key for i in block.members for key in self.computations[i].inputs} - outputs
|
|
origins = set().union(*(sources[key] for key in inputs)) if inputs else set()
|
|
for key in outputs:
|
|
sources[key] = origins
|
|
result.append({
|
|
'cyclic': block.cyclic,
|
|
'inputs': sorted(map(self.label, inputs)),
|
|
'outputs': sorted(map(self.label, outputs)),
|
|
'origins': sorted(origins),
|
|
'operations': [self.computations[i].key for i in block.members],
|
|
'pressureUnknowns': [self.label(key) for i in block.members
|
|
if self.computations[i].kind == 'pressure'
|
|
for key in self.computations[i].outputs],
|
|
})
|
|
return {
|
|
'strategy': 'dependency-blocks',
|
|
'operationCount': len(self.computations),
|
|
'cyclicBlockCount': sum(b.cyclic for b in self.blocks),
|
|
'knownSources': {self.label(key): value for key, value in self.known.items()},
|
|
'blocks': result,
|
|
'operations': [{'key': op.key, 'kind': op.kind,
|
|
'inputs': sorted(map(self.label, op.inputs)),
|
|
'outputs': list(map(self.label, op.outputs))}
|
|
for op in self.computations],
|
|
}
|
|
|
|
def emit(self):
|
|
"""Return C helper definitions and a straight-line/local-block schedule.
|
|
|
|
Existing scalar pressure bisection and stream convergence tolerances are
|
|
retained. Only the computations in the relevant SCC participate in each
|
|
closure; a pressure trial further restricts work to its residual inputs.
|
|
"""
|
|
helpers, lines = [], []
|
|
args = 'p,h,g,w,q,fb,pipe_cache'
|
|
signature = ('double *p,double *h,const NativeGas *g,double *w,double *q,'
|
|
'double *fb,NativePipeCache *pipe_cache')
|
|
|
|
def code(indices):
|
|
return [line for i in indices for line in
|
|
(f'/* schedule operation {i}: {self.computations[i].kind} */', *self.computations[i].code)]
|
|
|
|
def helper(name, indices):
|
|
helpers.extend([f'static int {name}({signature}) {{',
|
|
'(void)p;(void)h;(void)g;(void)w;(void)q;(void)fb;(void)pipe_cache;',
|
|
*code(indices), 'return 1;', '}'])
|
|
return f'if(!{name}({args})) return 0;'
|
|
|
|
for number, block in enumerate(self.blocks):
|
|
if not block.cyclic:
|
|
if self.computations[block.members[0]].kind == 'pressure':
|
|
raise NativeCapabilityError('Pressure balance has no pressure-dependent flow relation')
|
|
lines.extend(code(block.members))
|
|
continue
|
|
pressures = [i for i in block.members if self.computations[i].kind == 'pressure']
|
|
streams = [i for i in block.members if self.computations[i].kind in ('stream', 'alias')]
|
|
flows = set(block.members) - set(pressures) - set(streams)
|
|
if all(self.computations[i].kind == 'alias' for i in block.members):
|
|
raise NativeCapabilityError('Enthalpy reference cycle has no thermodynamic source: ' +
|
|
', '.join(self.computations[i].key for i in block.members))
|
|
refresh = helper(f'model_block_{number}_flows', self.ordered_subset(flows)) if flows else ''
|
|
lines.append(f'{{ /* local algebraic block {number} */')
|
|
hvars = [key for i in streams for key in self.computations[i].outputs]
|
|
outputs = {key for i in block.members for key in self.computations[i].outputs}
|
|
inputs = sorted({key for i in block.members for key in self.computations[i].inputs} - outputs)
|
|
if hvars:
|
|
seeds = [key for key in inputs if key.startswith('h[') or key.endswith('.h')]
|
|
if not seeds:
|
|
raise NativeCapabilityError('Local stream loop has no supplied thermodynamic state')
|
|
lines.extend([*[f'{key}={seeds[0]};' for key in hvars], 'int closure_ok=0;',
|
|
f'for(int closure=0;closure<{max(64,4*len(hvars))};closure++) {{',
|
|
'double previous[]={' + ','.join(hvars) + '};'])
|
|
if pressures:
|
|
pressure_bounds = [key for key in inputs if key.startswith('p[') or key.endswith('.p')]
|
|
if not pressure_bounds:
|
|
raise NativeCapabilityError('Local pressure block has no pressure boundary')
|
|
lines.extend(['double plo=INFINITY,phi=0;',
|
|
*[f'plo=fmin(plo,{p});phi=fmax(phi,{p});' for p in pressure_bounds],
|
|
*[f'{self.computations[i].outputs[0]}=.5*(plo+phi);' for i in pressures],
|
|
'int pressure_ok=0;', 'for(int sweep=0;sweep<256;sweep++) {'])
|
|
for i in pressures:
|
|
op = self.computations[i]
|
|
p = op.outputs[0]
|
|
trial = helper(f'model_block_{number}_pressure_{i}', self.ancestors(op.inputs, flows))
|
|
lines.extend(['{ double lo=plo,hi=phi;', 'for(int bisect=0;bisect<48;bisect++) {',
|
|
f'{p}=.5*(lo+hi);', trial,
|
|
f'double balance={op.residual};if(!isfinite(balance)) return 0;',
|
|
f'if(balance>0) hi={p};else lo={p};', '}}'])
|
|
lines.extend([refresh, 'double residual=0;',
|
|
*[f'{{double balance={self.computations[i].residual};if(!isfinite(balance)) return 0;residual=fmax(residual,fabs(balance));}}' for i in pressures],
|
|
'if(residual<1e-11) {pressure_ok=1;break;}', '}',
|
|
'if(!pressure_ok) return 0;'])
|
|
elif refresh:
|
|
lines.append(refresh)
|
|
if hvars:
|
|
# Gauss-Seidel within a genuine stream loop, in stable emission order.
|
|
lines.extend([*code(streams), 'double change=0;',
|
|
*[line for j,key in enumerate(hvars) for line in
|
|
(f'if(!isfinite({key})) return 0;',
|
|
f'change=fmax(change,fabs({key}-previous[{j}])/fmax(1,fabs({key})));')],
|
|
'if(change<1e-12) {closure_ok=1;break;}', '}',
|
|
'if(!closure_ok) return 0;', refresh])
|
|
lines.append('}')
|
|
return helpers, lines
|