From 9b4485ca289de67935491c007f58c7a543a93819 Mon Sep 17 00:00:00 2001 From: Claude Date: Fri, 18 Sep 2026 19:48:21 +0000 Subject: [PATCH] refactor(language): a partition is its own class rather than a direction with nothing consumed or produced A translation, a window and a grouped position hold a Partition: the key column stepped along, the group columns within= named, and the key columns joined on. Direction keeps sum and at, where columns are consumed and produced and the frame changes. A partition's frame does not change, so its fields no longer have to be read against a paragraph that redefines them. RelationNode carries either as `use`, and each consumer asserts the kind its operator takes. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01UnbTmmYp7auaiCwSgkKKsB --- src/math_spec/_expression_parser.py | 8 ++-- src/math_spec/dimensions.py | 18 ++++---- src/math_spec/exclusivity.py | 2 +- src/math_spec/lowering.py | 13 +++--- src/math_spec/program.py | 70 ++++++++++++++++++++++------- src/math_spec/resolution.py | 41 +++++++++-------- src/math_spec/typesetting/walk.py | 37 ++++++++------- tests/test_lowering.py | 7 +-- 8 files changed, 120 insertions(+), 76 deletions(-) diff --git a/src/math_spec/_expression_parser.py b/src/math_spec/_expression_parser.py index 1bd0dec9..c02958c4 100644 --- a/src/math_spec/_expression_parser.py +++ b/src/math_spec/_expression_parser.py @@ -22,7 +22,7 @@ if TYPE_CHECKING: from collections.abc import Callable, Iterator, Mapping - from math_spec.program import Direction, WhereNode + from math_spec.program import Direction, Partition, WhereNode #: The relation a comparison may carry — the three an expression may be #: written with, which is what a constraint's sense is read off. @@ -115,8 +115,10 @@ def shown(self) -> str: @dataclass(frozen=True) class RelationNode: - """A resolved ``by=`` — the relation, and the direction the call reads it in. + """A resolved ``by=`` — the relation, and the use the call makes of it. + ``use`` is a :class:`Direction` for ``sum`` and ``at``, and a + :class:`Partition` for ``shift`` and ``sum_back``. ``dimensions`` is the fine side — what ``sum`` consumes and ``at`` produces — and ``into`` the coarse dims, which ``sum`` produces and ``at`` consumes. The roles joined on are the operand's to carry, and the operator @@ -126,7 +128,7 @@ class RelationNode: name: str dimensions: tuple[str, ...] into: tuple[str, ...] - direction: Direction + use: Direction | Partition @property def shown(self) -> str: diff --git a/src/math_spec/dimensions.py b/src/math_spec/dimensions.py index 064c830c..b620db9b 100644 --- a/src/math_spec/dimensions.py +++ b/src/math_spec/dimensions.py @@ -45,6 +45,7 @@ Mask, ParameterComparisonNode, ParameterDefinedNode, + Partition, RelationComparisonNode, RelationDefinedNode, RelationPairComparisonNode, @@ -228,18 +229,18 @@ def _check_lands_clear(call: str, produced: set[str], consumed: set[str], inner: def _check_joined(call: str, by: RelationNode, inner: frozenset[str], context: str) -> None: """The columns a call joins on are read at their dimensions, so the operand carries every one, each once.""" - direction = by.direction - dims = direction.joined_dims + use = by.use + dims = use.joined_dims if missing := sorted(set(dims) - inner): raise DimensionError( - f'{context}: {call} joins on {missing} (columns {[r for r in direction.joined if direction.dim(r) in missing]} ' - f"of '{direction.name}'), which the expression does not carry (dims {sorted(inner)}). A relation is " + f'{context}: {call} joins on {missing} (columns {[r for r in use.joined if use.dim(r) in missing]} ' + f"of '{use.name}'), which the expression does not carry (dims {sorted(inner)}). A relation is " f'read between two of its columns and joined at the others — index the operand by them, or ' f'read it between different columns.' ) if twice := sorted({d for d in dims if dims.count(d) > 1 or d in by.dimensions}): raise DimensionError( - f"{context}: {call} joins '{direction.name}' on {twice} through more than one column, and the operand " + f"{context}: {call} joins '{use.name}' on {twice} through more than one column, and the operand " f'carries each dimension once. Read between different columns, or use a relation whose joined ' f'columns are over distinct dimensions.' ) @@ -421,11 +422,8 @@ def _check_named_amount(node: FunctionCallNode, over: str, inner: frozenset[str] f"— declare '{amount.name}' over dims '{over}' is not one of." ) partition = node.kwargs.get('by') - groups = ( - frozenset(partition.direction.dim(v) for v in partition.direction.produced) - if isinstance(partition, RelationNode) - else frozenset() - ) + use = partition.use if isinstance(partition, RelationNode) else None + groups = frozenset(use.dim(v) for v in use.group) if isinstance(use, Partition) else frozenset() if stray := sorted(frozenset(declared.dims) - inner - groups): raise DimensionError( f'{context}: {node.name}({kwarg}={amount.name}) reads its {words.noun} at the coordinate it ' diff --git a/src/math_spec/exclusivity.py b/src/math_spec/exclusivity.py index 75ddcb45..de1b652a 100644 --- a/src/math_spec/exclusivity.py +++ b/src/math_spec/exclusivity.py @@ -223,7 +223,7 @@ def _subject_of(node: TypedPredicateNode) -> Subject: case DimensionPositionNode(name=name, partition=partition): if partition is None: return Subject('rank', name) - return Subject('rank', name, partition.name, partition.produced) + return Subject('rank', name, partition.name, partition.group) case RelationDefinedNode(name=name) | RelationComparisonNode(name=name): return Subject('relation', name) case RelationPairComparisonNode(name=name, other=other): diff --git a/src/math_spec/lowering.py b/src/math_spec/lowering.py index 6f931266..9392a7db 100644 --- a/src/math_spec/lowering.py +++ b/src/math_spec/lowering.py @@ -274,13 +274,15 @@ def sum(self, node: FunctionCallNode) -> program.ExpressionNode: assert isinstance(consumed, DimensionNode), 'resolution refuses a over= that is not a dimension' return program.Sum(operand, (consumed.name,)) assert isinstance(by_node, RelationNode), 'resolution refuses a by= that is not a relation' - return program.GroupSum(operand, direction=by_node.direction) + assert isinstance(by_node.use, program.Direction), 'resolution reads sum(by=) in a direction' + return program.GroupSum(operand, direction=by_node.use) def at(self, node: FunctionCallNode) -> program.ExpressionNode: """``at(x, by=relation)`` — the adjoint of :meth:`sum`'s ``by=`` form.""" by_node = node.kwargs['by'] assert isinstance(by_node, RelationNode), 'resolution refuses a by= that is not a relation' - return program.At(self.expr(node.args[0]), direction=by_node.direction) + assert isinstance(by_node.use, program.Direction), 'resolution reads at(by=) in a direction' + return program.At(self.expr(node.args[0]), direction=by_node.use) def sum_back(self, node: FunctionCallNode) -> program.ExpressionNode: """``sum_back(x, along=d, window=w)`` — a trailing window along one dimension. @@ -342,8 +344,8 @@ def shift(self, node: FunctionCallNode) -> program.ExpressionNode: } -def _partition_of(node: FunctionCallNode) -> program.Direction | None: - """The direction a translation partitions by, if the call names a relation. +def _partition_of(node: FunctionCallNode) -> program.Partition | None: + """The partition a translation steps inside, if the call names a relation. That it is a *single* relation, stepped *along the translated dimension*, is checked with the other dim rules (``math_spec.dimensions``), where a model @@ -353,7 +355,8 @@ def _partition_of(node: FunctionCallNode) -> program.Direction | None: if by_node is None: return None assert isinstance(by_node, RelationNode) - return by_node.direction + assert isinstance(by_node.use, program.Partition), "resolution reads a translation's by= as a partition" + return by_node.use def _bound_expression(value: float | str) -> program.ExpressionNode: diff --git a/src/math_spec/program.py b/src/math_spec/program.py index 6d76a5c3..c9c80ffe 100644 --- a/src/math_spec/program.py +++ b/src/math_spec/program.py @@ -81,6 +81,7 @@ 'ParameterDeclaration', 'ParameterDefinedNode', 'ParameterDtype', + 'Partition', 'PiecewiseDeclaration', 'Power', 'PredicateOperator', @@ -342,12 +343,11 @@ class Translate(Expression): ``offset`` is an integer, or the name of an integer parameter that does not depend on ``dimension`` and carries its sign in the values. - ``partition`` is a relation stepped along ``dimension`` — its consumed - column is a key over that dimension, its produced columns are the group — - and the translation then happens inside each group: the neighbour is the - one before in the same group, the edge is the group's, and a wrap closes - each group onto itself. A coordinate the relation sends nowhere reaches - nothing. + ``partition`` is a relation with a key column over ``dimension`` + (:class:`Partition`), and the translation then happens inside each group + its ``within=`` columns make: the neighbour is the one before in the same + group, the edge is the group's, and a wrap closes each group onto itself. + A coordinate the relation sends nowhere reaches nothing. """ operand: ExpressionNode @@ -355,7 +355,7 @@ class Translate(Expression): offset: int | str wrap: bool fill: float | None = None - partition: Direction | None = None + partition: Partition | None = None @dataclass(frozen=True) @@ -383,7 +383,7 @@ class Window(Expression): dimension: str width: int | str wrap: bool - partition: Direction | None = None + partition: Partition | None = None @dataclass(frozen=True) @@ -519,9 +519,7 @@ class Direction(NamedTuple): of ``relation``, which binds every role to its dimension and names the key. ``joined`` is the key roles the call did not name (every role, for a bare relation): the join keys on them, and a value role left unnamed is not - read. For a partition (``shift``, ``sum_back``, ``position``) ``consumed`` - is the key role over the dimension stepped along and ``produced`` the value roles - that make the group, which are the ones ``within=`` named. + read. """ relation: RelationDeclaration @@ -567,6 +565,48 @@ def is_function_read(self) -> bool: return bool(self.key) and set(self.key) <= {*self.joined, *self.produced} +class Partition(NamedTuple): + """One relation as a partition steps along it — the key column stepped along, the group columns, and the key columns joined on. + + ``along``, ``group`` and ``joined`` are *roles* — column names of + ``relation``, which binds every role to its dimension and names the key. + ``along`` is the one key column over the dimension stepped along, and + the frame keeps it. ``group`` is the value columns ``within=`` named, + read at the row's key. ``joined`` is the other key columns, whose + dimensions the frame carries. Nothing is consumed and nothing is + produced: the frame does not change. + """ + + relation: RelationDeclaration + along: str + group: tuple[str, ...] + joined: tuple[str, ...] + + @property + def name(self) -> str: + return self.relation.name + + @property + def key(self) -> tuple[str, ...]: + return self.relation.key + + @property + def values(self) -> tuple[str, ...]: + return self.relation.values + + def dim(self, role: str) -> str: + """The dimension *role* is bound to.""" + return self.relation.dim(role) + + @property + def along_dim(self) -> str: + return self.dim(self.along) + + @property + def joined_dims(self) -> tuple[str, ...]: + return tuple(self.dim(role) for role in self.joined) + + @dataclass(frozen=True) class DimensionDeclaration: """A dimension and the relations with a column over it.""" @@ -1212,16 +1252,14 @@ class DimensionPositionNode: """Compare where a row sits along a dimension against a position — ``position(snapshot) == 0``. Both sides are integers, negative counting from the end. With a - ``partition`` the position is counted within each group the relation makes, - read as :class:`Translate` reads one: its consumed column is the key - column over ``name``, the group is its produced columns, and its joined - columns are the other key columns, whose dimensions the frame carries. + ``partition`` the position is counted within each group the relation makes + (:class:`Partition`), whose joined columns' dimensions the frame carries. """ name: str op: PredicateOperator position: int - partition: Direction | None = None + partition: Partition | None = None @dataclass(frozen=True) diff --git a/src/math_spec/resolution.py b/src/math_spec/resolution.py index 741bca02..adeaf699 100644 --- a/src/math_spec/resolution.py +++ b/src/math_spec/resolution.py @@ -72,6 +72,7 @@ OrNode, ParameterComparisonNode, ParameterDefinedNode, + Partition, PredicateOperator, RelationComparisonNode, RelationDeclaration, @@ -609,20 +610,18 @@ def _relation_ref( if 'within' not in named: return value # the call shape refused it already, with the wording that names the rewrite over_dim = over.name if isinstance(over, NameNode | DimensionNode) else None - direction = self._partition_direction(name, operator, over_dim, named['within']) - else: - if not ({'over', 'into'} <= set(named)): - return value # the call shape refused it already, with the wording that names the rewrite - direction = self._direction(name, operator, named['over'], named['into']) + partition = self._partition(name, operator, over_dim, named['within']) + if partition is None: + return value + return RelationNode(name, dimensions=(partition.along_dim,), into=(), use=partition) + if not ({'over', 'into'} <= set(named)): + return value # the call shape refused it already, with the wording that names the rewrite + direction = self._direction(name, operator, named['over'], named['into']) if direction is None: return value - fine = direction.produced_dims if operator == 'at' else direction.consumed_dims - if operator in ('shift', 'sum_back'): - coarse: tuple[str, ...] = () - else: - coarse = direction.consumed_dims if operator == 'at' else direction.produced_dims - return RelationNode(name, dimensions=fine, into=coarse, direction=direction) + coarse = direction.consumed_dims if operator == 'at' else direction.produced_dims + return RelationNode(name, dimensions=fine, into=coarse, use=direction) def _role_name(self, value: ArithmeticNode, operator: str, key: str) -> tuple[str, ...] | None: """``over=`` or ``into=`` as the column names it must be — one bare name, or a bracketed list of them.""" @@ -708,14 +707,14 @@ def _known_roles(self, name: str, call: str, roles: tuple[str, ...] | None, kwar return False return True - def _partition_direction( + def _partition( self, name: str, operator: str, along_dim: str | None, within_roles: tuple[str, ...] - ) -> Direction | None: - """Which direction a partition (``shift``, ``sum_back``, ``position``) reads relation *name* in along *along_dim*. + ) -> Partition | None: + """How a partition (``shift``, ``sum_back``, ``position``) steps along relation *name* over *along_dim*. - It takes the one key column over that dimension (a key has one column - per dimension), joins on the other key columns and groups by the value - columns *within_roles* names. ``None`` where the dimension is not one + It steps along the one key column over that dimension (a key has one + column per dimension), joins on the other key columns and groups by the + value columns *within_roles* names. ``None`` where the dimension is not one (already refused), the relation has no key column over it, or ``within=`` names a column that is not a value column. """ @@ -746,7 +745,7 @@ def _partition_direction( return None (along,) = over_keys joined = tuple(r for r in shape.key if r != along) - return Direction(shape, (along,), within_roles, joined) + return Partition(shape, along, within_roles, joined) def _not_a_relation(self, name: str, operator: str, key: str) -> str | None: """Why *name* is not a relation; ``None`` where it is one.""" @@ -860,10 +859,10 @@ def _position(self, node: UnresolvedPositionNode) -> DimensionPositionNode | Unr f'are {list(ns.shape_of(node.by).values)}.' ) return node - direction = self._partition_direction(node.by, 'position', node.dimension, node.into) - if direction is None: + partition = self._partition(node.by, 'position', node.dimension, node.into) + if partition is None: return node - return DimensionPositionNode(node.dimension, node.op, node.position, direction) + return DimensionPositionNode(node.dimension, node.op, node.position, partition) def _comparison(self, node: UnresolvedComparisonNode) -> WhereNode | UnresolvedWhereNode: """``name literal``, or the one structural form ``relation relation``.""" diff --git a/src/math_spec/typesetting/walk.py b/src/math_spec/typesetting/walk.py index b2e40177..6123b74c 100644 --- a/src/math_spec/typesetting/walk.py +++ b/src/math_spec/typesetting/walk.py @@ -39,11 +39,13 @@ BooleanLiteralNode, DimensionComparisonNode, DimensionPositionNode, + Direction, Mask, NotNode, OrNode, ParameterComparisonNode, ParameterDefinedNode, + Partition, PredicateOperator, RelationComparisonNode, RelationDefinedNode, @@ -58,7 +60,7 @@ from collections.abc import Iterable, Mapping from math_spec.model import RelationBlock, SosBlock, _ExpandedSpec - from math_spec.program import Direction + from math_spec.program import RelationDeclaration from math_spec.typesetting.format import Format from math_spec.typesetting.symbols import Symbols @@ -274,15 +276,15 @@ def _translation(self, step: _Step) -> str: self.noticed.grouped = True return self.format.superscript(operator, step.within) - def _relation_read(self, direction: Direction, at: Mapping[str, str], read: str) -> str: + def _relation_read(self, relation: RelationDeclaration, at: Mapping[str, str], read: str) -> str: """A relation's column *read* as a function at the columns *at* fixes: ``bus(g)``, ``zone_of(g, p)`` or ``ends.bus0(l)``. *at* maps each key role to the index it is read at. The function is named after the relation alone where the key determines one column, and after the column read otherwise. """ - name = direction.name if len(direction.values) == 1 else f'{direction.name}.{read}' - return self.format.apply(self.format.upright(name), self.format.joined([at[k] for k in direction.key], '')) + name = relation.name if len(relation.values) == 1 else f'{relation.name}.{read}' + return self.format.apply(self.format.upright(name), self.format.joined([at[k] for k in relation.key], '')) def _relation_row(self, name: str, key: list[str]) -> str: """That relation *name* has a row at *key*, its key columns' indices in declared order. @@ -307,12 +309,10 @@ def _value_read(self, name: str, column: str, ctx: _Context) -> str: def _position_group(self, node: DimensionPositionNode, ctx: _Context) -> str: """The group a grouped position counts within: the relation's group columns read at the row's key.""" assert node.partition is not None - direction = node.partition - keyed = self.format.joined([ctx.subscript(direction.dim(k)) for k in direction.key], '') - single = len(direction.values) == 1 - reads = [ - self.format.apply(self._column(direction.name, column, single), keyed) for column in direction.produced - ] + partition = node.partition + keyed = self.format.joined([ctx.subscript(partition.dim(k)) for k in partition.key], '') + single = len(partition.values) == 1 + reads = [self.format.apply(self._column(partition.name, column, single), keyed) for column in partition.group] return self._tuple(reads) def _tuple(self, reads: list[str]) -> str: @@ -459,10 +459,11 @@ def _call(self, node: FunctionCallNode, ctx: _Context) -> tuple[str, int]: by = node.kwargs['by'] assert isinstance(by, RelationNode) outer = ctx - direction = by.direction + direction = by.use + assert isinstance(direction, Direction) at = {r: outer.subscript(direction.dim(r)) for r in (*direction.produced, *direction.joined)} for read in direction.consumed: - ctx = ctx.pulled_back(direction.dim(read), self._relation_read(direction, at, read)) + ctx = ctx.pulled_back(direction.dim(read), self._relation_read(direction.relation, at, read)) return self._arithmetic(node.args[0], ctx) if (by := node.kwargs.get('by')) is not None: @@ -471,7 +472,8 @@ def _call(self, node: FunctionCallNode, ctx: _Context) -> tuple[str, int]: inner = ctx for d in by.dimensions: dummies[d], inner = inner.reducing(d) - conditions = list(self._grouping(by.direction, dummies, ctx)) + assert isinstance(by.use, Direction) + conditions = list(self._grouping(by.use, dummies, ctx)) domain = ( f'{self.format.joined([self._membership(d, dummies[d]) for d in by.dimensions], "")} ' f'{self._op("such_that")} {self.format.joined(conditions, self._op("and"))}' @@ -505,7 +507,7 @@ def _grouping(self, direction: Direction, dummies: Mapping[str, str], ctx: _Cont fixed = [r for r in direction.values if r in at] if not fixed: return [self._relation_row(direction.name, [at[k] for k in direction.key])] - return [f'{self._relation_read(direction, at, r)} {self._op("equal")} {at[r]}' for r in fixed] + return [f'{self._relation_read(direction.relation, at, r)} {self._op("equal")} {at[r]}' for r in fixed] def _group(self, by: ArithmeticNode | None, dim: str) -> str: """A ``by=`` as the superscript its translation operator carries. @@ -517,9 +519,10 @@ def _group(self, by: ArithmeticNode | None, dim: str) -> str: if by is None: return '' assert isinstance(by, RelationNode) - direction = by.direction - at = {r: self.symbols.index[direction.dim(r)] for r in (*direction.consumed, *direction.joined)} - return self._tuple([self._relation_read(direction, at, r) for r in direction.produced]) + partition = by.use + assert isinstance(partition, Partition) + at = {r: self.symbols.index[partition.dim(r)] for r in (partition.along, *partition.joined)} + return self._tuple([self._relation_read(partition.relation, at, r) for r in partition.group]) def _width(self, node: ArithmeticNode) -> str: """``sum_back``'s ``window=``: a number, or a parameter's own symbol. diff --git a/tests/test_lowering.py b/tests/test_lowering.py index 58f85e3f..31e5b47b 100644 --- a/tests/test_lowering.py +++ b/tests/test_lowering.py @@ -45,6 +45,7 @@ Parameter, ParameterComparisonNode, ParameterDefinedNode, + Partition, Power, Program, Region, @@ -443,7 +444,7 @@ def test_a_power_lowers_to_a_node_of_its_own(dispatch_schema): offset=1, wrap=False, fill=0.0, - partition=Direction(LK, ('g',), ('h',), ()), + partition=Partition(LK, 'g', ('h',), ()), ), id='a-translation-stops-at-the-edges-of-the-relation-it-names', ), @@ -464,7 +465,7 @@ def test_a_power_lowers_to_a_node_of_its_own(dispatch_schema): 'g', width=2, wrap=False, - partition=Direction(LK, ('g',), ('h',), ()), + partition=Partition(LK, 'g', ('h',), ()), ), id='a-window-stops-at-the-edges-of-the-relation-it-names', ), @@ -498,7 +499,7 @@ def test_a_partition_keeps_its_group_when_the_relation_gains_a_value_column(): }, } ) - grouping[str(values)] = _partition_of(program.constraints['k']).produced + grouping[str(values)] = _partition_of(program.constraints['k']).group assert grouping == {'day': ('day',), "['day', 'week']": ('day',)}, ( 'the group is the columns the call named, on both calendars' )