Skip to content

math_spec.separability

The walk behind :attr:~math_spec.program.Program.separability — every axis's verdict, in one pass over a program.

separabilities(program) #

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.

Source code in src/math_spec/separability.py
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
    }