def separabilities(program: Program) -> dict[str, Separability]:
"""Every axis's verdict, in one walk.
One traversal rather than one per axis, because every construct that ties an
axis together names the axis it ties: asking each node *which* dimension it
is about answers for all of them at what answering for one cost.
``reductions_couple`` is the position a block stands in rather than anything
about the block — a sum over the axis couples a constraint row to the whole
horizon and leaves an objective additively separable. A translation reads
ahead for a negative offset; what one reads behind is the window's edge,
which is not asked. Each coupling carries the one modelling change that
would lift it, after the dash.
The border a decomposition cuts along falls out of the same walk, which is
why it is taken here rather than in a pass of its own: a constraint the axis
does not index, and one a coupling names, are the rows no window holds.
"""
ahead = dict.fromkeys(program.dimensions, 0)
reasons: dict[str, dict[str, dict[str, list[str]]]] = {
kind: {dimension: {} for dimension in program.dimensions} for kind in ('coupled', 'restarts')
}
undecided: dict[str, dict[Reach, None]] = {dimension: {} for dimension in program.dimensions}
rows: dict[str, str] = {}
def report(kind: str, dimension: str, label: str, reason: str) -> None:
reasons[kind][dimension].setdefault(label, []).append(reason)
def waits_on(dimension: str, label: str, name: str, kind: Literal['offset', 'partition', 'coordinate']) -> None:
undecided[dimension][Reach(label, name, kind)] = None
for label, row, nodes, mask, reductions_couple in _built_blocks(program):
if row is not None:
rows[label] = row
masks: list[Mask | None] = [mask]
for node in walk(*nodes):
if isinstance(node, Cases):
masks.extend(region.when for region in node.regions)
elif isinstance(node, Sum):
if reductions_couple:
for dimension in node.over:
report(
'coupled',
dimension,
label,
f'sums over {dimension} — a rolling sum_back(window=n) windows, a total over the horizon does not',
)
elif isinstance(node, GroupSum):
for dimension in node.direction.consumed_dims:
report(
'coupled',
dimension,
label,
f'groups {dimension} into {", ".join(node.direction.produced_dims)} — window that dimension instead, or cut only at the group edges',
)
elif isinstance(node, Pullback):
for dimension in node.direction.consumed_dims:
waits_on(dimension, label, node.direction.name, 'coordinate')
elif isinstance(node, (Translate, WindowSum)):
dimension = node.along
if node.wrap:
report(
'coupled',
dimension,
label,
f'wraps around {dimension}, so its first row reads its last — an opening-state seed at '
f'position({dimension}) == 0 is what a rolling horizon replaces the wrap with',
)
continue
if node.partition is not None:
waits_on(dimension, label, node.partition.name, 'partition')
if isinstance(node, WindowSum):
continue
if isinstance(node.offset, str):
waits_on(dimension, label, node.offset, 'offset')
else:
ahead[dimension] = max(ahead[dimension], -node.offset)
for candidate in masks:
for atom in candidate.atoms if candidate is not None else ():
if isinstance(atom, DimensionPosition):
report('restarts', atom.name, label, f'counts a position along {atom.name}')
for name, block in program.sos.items():
report(
'coupled',
block.along,
f"set '{name}'",
f'is a set along {block.along}, which a window would cut — only a window holding every whole set keeps it',
)
def joined(kind: str, dimension: str) -> dict[str, str]:
return {label: ', '.join(dict.fromkeys(found)) for label, found in reasons[kind][dimension].items()}
def linking_rows(dimension: str) -> tuple[str, ...]:
"""Each constraint no one window of *dimension* holds whole, in declaration order.
Two shapes reach the border by different routes, and a constraint that
takes both is still one name: a row the axis does not index stands in
every window, and a row a coupling names reads the whole axis. Only a
declaration that builds a row can put one here, which is what ``rows``
holds the coupled labels to.
"""
coupled = {rows[label] for label in reasons['coupled'][dimension] if label in rows}
return tuple(
name
for name, constraint in program.constraints.items()
if dimension not in constraint.dims or name in coupled
)
return {
dimension: Separability(
dimension=dimension,
ahead=ahead[dimension],
coupled=joined('coupled', dimension),
undecided=tuple(undecided[dimension]),
restarts=joined('restarts', dimension),
linking_rows=linking_rows(dimension),
linking_columns=tuple(
name for name, variable in program.variables.items() if dimension not in variable.dims
),
)
for dimension in program.dimensions
}