From 39c160677c3f8bd1946184278c083a5bbe99ddff Mon Sep 17 00:00:00 2001 From: Claude Date: Tue, 15 Sep 2026 10:04:21 +0000 Subject: [PATCH 1/2] refactor(program): a program names its nodes by the naming rule and its groups as the file does The where vocabulary drops the parser's `Node` suffix, so `And`, `Not`, `Or` and `ParameterComparison` stand beside `Sum` and `Add` as the naming rule says; the unions are `Expression`, `Predicate`, `TypedPredicate` and `Connective`. `At` is `Pullback`, named for the coordinate map rather than the file's spelling, and `Window` is `WindowSum`, paired with `GroupSum`. A translation and a window call their axis `over`, as every other node does. The file's `expressions:` section is `Program.expressions`, and the trees a row is built from are `Program.roots`. Lookups are one group keyed by name, each with `over` and `into`, as the file declares them, rather than nested under a dimension with `target`. `Footprint.shapes` is `nodes`. Gone: the base class and its `+` and `*` sugar, the singular accessors `dimension()`, `parameter()` and `variable()`, and `DimensionDeclaration.targets`. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_019fhGZgaBspo7mh9Hjd3KtT --- docs/contributing.md | 10 +- docs/reference/language/reading.md | 4 +- docs/reference/language/reported.md | 2 +- src/math_spec/_expression_parser.py | 6 +- src/math_spec/_where_parser.py | 24 +- src/math_spec/advice.py | 8 +- src/math_spec/boundedness.py | 12 +- src/math_spec/dimensions.py | 24 +- src/math_spec/exclusivity.py | 80 +++--- src/math_spec/lowering.py | 32 +-- src/math_spec/program.py | 408 ++++++++++++---------------- src/math_spec/resolution.py | 88 +++--- src/math_spec/separability.py | 20 +- src/math_spec/typesetting/walk.py | 54 ++-- src/math_spec/validation.py | 8 +- tests/test_dimensions.py | 4 +- tests/test_exclusivity.py | 12 +- tests/test_lowering.py | 205 +++++++------- tests/test_parser.py | 28 +- tests/test_piecewise.py | 2 +- tests/test_program_nodes.py | 12 +- tests/test_validation.py | 6 +- tests/typesetting/test_golden.py | 4 +- 23 files changed, 487 insertions(+), 566 deletions(-) diff --git a/docs/contributing.md b/docs/contributing.md index 89815cdd..362126d3 100644 --- a/docs/contributing.md +++ b/docs/contributing.md @@ -126,11 +126,11 @@ stale anchor fails it. The same construct passes through three layers, and each names it in full. The suffix says which layer, which keeps the three vocabularies from colliding: -| Layer | Suffix | Example | -| ------------------------------- | -------------------- | ----------------------------------------- | -| YAML block (`math_spec.model`) | `Block` | `VariableBlock`, `PiecewiseBlock` | -| Core AST (`math_spec.*_parser`) | `Node` | `VariableNode`, `DimensionComparisonNode` | -| Program (`math_spec.program`) | none / `Declaration` | `Variable`, `VariableDeclaration` | +| Layer | Suffix | Example | +| ------------------------------- | -------------------- | ------------------------------------------ | +| YAML block (`math_spec.model`) | `Block` | `VariableBlock`, `PiecewiseBlock` | +| Core AST (`math_spec.*_parser`) | `Node` | `VariableNode`, `UnresolvedComparisonNode` | +| Program (`math_spec.program`) | none / `Declaration` | `Variable`, `VariableDeclaration` | Two rules follow, and a PR that adds a construct keeps them: diff --git a/docs/reference/language/reading.md b/docs/reference/language/reading.md index 3b6abcc6..a01850ea 100644 --- a/docs/reference/language/reading.md +++ b/docs/reference/language/reading.md @@ -128,7 +128,7 @@ footprint = program.footprint sorted(footprint.quadratic) # [] sorted(footprint.domains) # ['continuous'] sorted(footprint.sos_types) # [] -sorted(kind.__name__ for kind in footprint.shapes) # ['Constant', 'Multiply', 'Parameter', 'Sum', 'Variable'] +sorted(kind.__name__ for kind in footprint.nodes) # ['Constant', 'Multiply', 'Parameter', 'Sum', 'Variable'] ``` Every field is a set. `if footprint.sos_types` asks whether sets appear at all, @@ -143,7 +143,7 @@ model does not use the construct, not that the construct does not exist. on the numbers. The footprint stops at the kind of construct. An engine whose solver accepts a -window but not a wrapped one reads `Window in footprint.shapes`, then walks the +window but not a wrapped one reads `WindowSum in footprint.nodes`, then walks the tree for the detail. ## Asking whether an axis can be cut diff --git a/docs/reference/language/reported.md b/docs/reference/language/reported.md index 6d6fd1a9..178befb4 100644 --- a/docs/reference/language/reported.md +++ b/docs/reference/language/reported.md @@ -44,7 +44,7 @@ variable in it, such as `(1 + rate) ** period`, is reported all the same. Deciding by use costs one thing: an entry meant for a constraint, and never named there, loads as a reported quantity instead of failing. -An engine reads the answer at `Program.named_expressions[name].in_math`. +An engine reads the answer at `Program.expressions[name].in_math`. ## Which restrictions do not apply diff --git a/src/math_spec/_expression_parser.py b/src/math_spec/_expression_parser.py index 642a9f6a..e31990aa 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 WhereNode + from math_spec.program import Predicate #: The relation a comparison may carry — the three an expression may be #: written with, which is what a constraint's sense is read off. @@ -187,7 +187,7 @@ class CaseArm: """ label: str - when: WhereNode | None + when: Predicate | None value: ArithmeticNode @@ -259,7 +259,7 @@ class ComparisonNode: #: A whole spec-side expression tree — parse output and the resolved tree alike. -#: Named apart from :data:`math_spec.program.ExpressionNode`, the lowered +#: Named apart from :data:`math_spec.program.Expression`, the lowered #: vocabulary a consumer reads. ParsedNode = ArithmeticNode | ComparisonNode diff --git a/src/math_spec/_where_parser.py b/src/math_spec/_where_parser.py index ea873445..1a143c52 100644 --- a/src/math_spec/_where_parser.py +++ b/src/math_spec/_where_parser.py @@ -16,12 +16,12 @@ import pyparsing as pp from math_spec._expression_parser import NAME, REAL, parse_text -from math_spec.program import AndNode, BooleanLiteralNode, NotNode, OrNode, PredicateOperator, where_children +from math_spec.program import And, BooleanLiteral, Not, Or, PredicateOperator, where_children if TYPE_CHECKING: from collections.abc import Callable - from math_spec.program import WhereNode + from math_spec.program import Predicate # --------------------------------------------------------------------------- # AST nodes @@ -98,8 +98,8 @@ def _build_where_grammar() -> pp.ParserElement: """ where_expr = pp.Forward() - true_lit = pp.CaselessKeyword('True').set_parse_action(lambda: BooleanLiteralNode(True)) - false_lit = pp.CaselessKeyword('False').set_parse_action(lambda: BooleanLiteralNode(False)) + true_lit = pp.CaselessKeyword('True').set_parse_action(lambda: BooleanLiteral(True)) + false_lit = pp.CaselessKeyword('False').set_parse_action(lambda: BooleanLiteral(False)) # pyrefly: ignore[implicit-any-lambda] number = pp.Regex(rf'-?({REAL}|\d+)').set_parse_action(lambda t: float(t[0])) @@ -135,28 +135,28 @@ def _build_where_grammar() -> pp.ParserElement: NOT = pp.CaselessKeyword('NOT').suppress() # pyrefly: ignore[implicit-any-lambda] - not_expr = (NOT + atom).set_parse_action(lambda t: NotNode(t[0])) | atom + not_expr = (NOT + atom).set_parse_action(lambda t: Not(t[0])) | atom AND = pp.CaselessKeyword('AND').suppress() and_expr = not_expr + pp.ZeroOrMore(AND + not_expr) - and_expr.set_parse_action(_folder(AndNode)) + and_expr.set_parse_action(_folder(And)) OR = pp.CaselessKeyword('OR').suppress() or_expr = and_expr + pp.ZeroOrMore(OR + and_expr) - or_expr.set_parse_action(_folder(OrNode)) + or_expr.set_parse_action(_folder(Or)) where_expr <<= or_expr return where_expr -def _folder(node_type: type[AndNode] | type[OrNode]) -> Callable[[pp.ParseResults], Any]: +def _folder(node_type: type[And] | type[Or]) -> Callable[[pp.ParseResults], Any]: """A parse action left-folding a flat operator chain into *node_type*.""" def fold(tokens: pp.ParseResults) -> Any: items = list(tokens) - result: WhereNode | UnresolvedWhereNode = items[0] + result: Predicate | UnresolvedWhereNode = items[0] for item in items[1:]: - result = node_type(cast('WhereNode', result), item) + result = node_type(cast('Predicate', result), item) return result return fold @@ -193,7 +193,7 @@ def _named_rewrite(text: str, loc: int) -> str | None: @lru_cache(maxsize=4096) -def parse_where(text: str) -> WhereNode | UnresolvedWhereNode: +def parse_where(text: str) -> Predicate | UnresolvedWhereNode: """Parse a where string into an AST, its leaves still unresolved. The connectives and literals are the resolved vocabulary's own; the leaves @@ -207,6 +207,6 @@ def parse_where(text: str) -> WhereNode | UnresolvedWhereNode: complaint. """ return cast( - 'WhereNode | UnresolvedWhereNode', + 'Predicate | UnresolvedWhereNode', parse_text(_WHERE_GRAMMAR, text, 'where string', _named_rewrite, where_children, _DEEP_REWRITE), ) diff --git a/src/math_spec/advice.py b/src/math_spec/advice.py index 9e90b5a6..ec288343 100644 --- a/src/math_spec/advice.py +++ b/src/math_spec/advice.py @@ -14,7 +14,7 @@ from math_spec.boundedness import unbounded_notes from math_spec.errors import Advice from math_spec.lowering import to_program -from math_spec.program import At, GroupSum, walk +from math_spec.program import GroupSum, Pullback, walk if TYPE_CHECKING: from pathlib import Path @@ -51,7 +51,7 @@ def _never_an_axis(program: Program) -> list[Advice]: for declaration in (*program.parameters.values(), *program.variables.values(), *program.constraints.values()): reached.update(declaration.dims) reached |= _produced_axes(program) - reached |= {lk.target for _, lk in program.lookups} + reached |= {lk.into for lk in program.lookups.values()} return [ Advice( @@ -72,9 +72,9 @@ def _produced_axes(program: Program) -> set[str]: ``sum(by=)`` lands on its target and ``at()`` spreads onto its fine dimension. """ axes: set[str] = set() - for node in walk(*program.expressions): + for node in walk(*program.roots): if isinstance(node, GroupSum): axes.update(node.into) - elif isinstance(node, At): + elif isinstance(node, Pullback): axes.add(node.over) return axes diff --git a/src/math_spec/boundedness.py b/src/math_spec/boundedness.py index 0c6b3e7e..1102abe8 100644 --- a/src/math_spec/boundedness.py +++ b/src/math_spec/boundedness.py @@ -19,21 +19,21 @@ from math_spec.errors import Advice from math_spec.program import ( Add, - At, Cases, Constant, Divide, Dual, - ExpressionNode, + Expression, GroupSum, Multiply, Negate, Parameter, Power, + Pullback, Sum, Translate, Variable, - Window, + WindowSum, children, variables_of, ) @@ -113,7 +113,7 @@ def _times(sign: Sign, other: Sign) -> Sign: return None if sign is None or other is None else ('+' if sign == other else '-') -def _coefficient_sign(node: ExpressionNode) -> Sign: +def _coefficient_sign(node: Expression) -> Sign: """The sign *node* scales a term by, or ``None`` unless it is a signed constant. ``-2`` lowers to a negation over a constant, so the sign of a literal @@ -128,7 +128,7 @@ def _coefficient_sign(node: ExpressionNode) -> Sign: return None -def _record_signs(node: ExpressionNode, sign: Sign, signs: dict[str, Sign]) -> None: +def _record_signs(node: Expression, sign: Sign, signs: dict[str, Sign]) -> None: """Record the sign each variable under *node* carries into the objective. A variable reached twice with different signs, or once with an undecidable @@ -162,7 +162,7 @@ def _record_signs(node: ExpressionNode, sign: Sign, signs: dict[str, Sign]) -> N _record_signs(node.base, None, signs) _record_signs(node.exponent, None, signs) return - if isinstance(node, Sum | GroupSum | At | Translate | Window | Cases): + if isinstance(node, Sum | GroupSum | Pullback | Translate | WindowSum | Cases): for child in children(node): _record_signs(child, sign, signs) return diff --git a/src/math_spec/dimensions.py b/src/math_spec/dimensions.py index 3c07aed5..84d5945b 100644 --- a/src/math_spec/dimensions.py +++ b/src/math_spec/dimensions.py @@ -40,15 +40,15 @@ from math_spec.errors import DimensionError from math_spec.operators import BUILTINS from math_spec.program import ( - DimensionComparisonNode, - DimensionPositionNode, - LookupComparisonNode, - LookupDefinedNode, - LookupPairComparisonNode, + DimensionComparison, + DimensionPosition, + LookupComparison, + LookupDefined, + LookupPairComparison, Mask, - ParameterComparisonNode, - ParameterDefinedNode, - VariableDefinedNode, + ParameterComparison, + ParameterDefined, + VariableDefined, ) if TYPE_CHECKING: @@ -514,13 +514,13 @@ def _check_where_dims( if not (outside := sorted(Mask(atom).dims - frame)): continue match atom: - case ParameterDefinedNode() | ParameterComparisonNode(): + case ParameterDefined() | ParameterComparison(): noun = 'parameter' - case VariableDefinedNode(): + case VariableDefined(): noun = 'variable' - case DimensionComparisonNode() | DimensionPositionNode(): + case DimensionComparison() | DimensionPosition(): noun = 'dimension' - case LookupComparisonNode() | LookupPairComparisonNode() | LookupDefinedNode(): + case LookupComparison() | LookupPairComparison() | LookupDefined(): noun = 'lookup' case _: assert_never(atom) diff --git a/src/math_spec/exclusivity.py b/src/math_spec/exclusivity.py index 450aad92..1529cc88 100644 --- a/src/math_spec/exclusivity.py +++ b/src/math_spec/exclusivity.py @@ -22,27 +22,27 @@ from typing import TYPE_CHECKING, Any, Literal, assert_never, cast from math_spec.program import ( - AndNode, - BooleanLiteralNode, - DimensionComparisonNode, - DimensionPositionNode, - LookupComparisonNode, - LookupDefinedNode, - LookupPairComparisonNode, + And, + BooleanLiteral, + DimensionComparison, + DimensionPosition, + LookupComparison, + LookupDefined, + LookupPairComparison, Mask, - NotNode, - OrNode, - ParameterComparisonNode, - ParameterDefinedNode, - TypedPredicateNode, - VariableDefinedNode, + Not, + Or, + ParameterComparison, + ParameterDefined, + TypedPredicate, + VariableDefined, ) if TYPE_CHECKING: from collections.abc import Iterable, Iterator, Mapping from math_spec.model import DeclaredDtype - from math_spec.program import PredicateOperator, WhereNode + from math_spec.program import Predicate, PredicateOperator #: The most cells one pair may multiply out to; a pair past it is several expressions. CELL_BUDGET = 8192 @@ -56,7 +56,7 @@ class Undecidable(Exception): # noqa: N818 """A pair this procedure will not reason about. Carries the rewrite.""" -def overlapping(cases: Mapping[str, WhereNode], dtypes: Mapping[str, DeclaredDtype]) -> Iterator[str]: +def overlapping(cases: Mapping[str, Predicate], dtypes: Mapping[str, DeclaredDtype]) -> Iterator[str]: """One refusal per pair of cases that could both claim a coordinate. Args: @@ -89,7 +89,7 @@ def overlapping(cases: Mapping[str, WhereNode], dtypes: Mapping[str, DeclaredDty ) -def _witness(first: WhereNode, second: WhereNode, dtypes: Mapping[str, DeclaredDtype]) -> str | None: +def _witness(first: Predicate, second: Predicate, dtypes: Mapping[str, DeclaredDtype]) -> str | None: """A coordinate both masks claim, rendered — ``None`` where no cell holds both.""" masks = (Mask(first), Mask(second)) grid = _Grid.of(masks, dtypes) @@ -183,15 +183,15 @@ def witness(self, cell: dict[Subject, Cell]) -> str: return ', '.join(f'{subject} is {_shown(subject, value)}' for subject, value in cell.items()) -def _observe(node: TypedPredicateNode, subject: Subject, values: set[Any], dtypes: Mapping[str, DeclaredDtype]) -> None: +def _observe(node: TypedPredicate, subject: Subject, values: set[Any], dtypes: Mapping[str, DeclaredDtype]) -> None: """Record what *node* says about its subject: a position, or a literal. ``position()`` converts the dimension to an integer, so an ordering over a rank is an ordering of integers and every comparator is admitted there. """ - if isinstance(node, DimensionPositionNode): + if isinstance(node, DimensionPosition): values.add(node.position) - elif isinstance(node, LookupPairComparisonNode): + elif isinstance(node, LookupPairComparison): if node.op not in ('==', '!='): msg = ( f'{subject} is ordered with {node.op!r}, and two lookups carry no order ' @@ -199,7 +199,7 @@ def _observe(node: TypedPredicateNode, subject: Subject, values: set[Any], dtype f'ordering as a boolean parameter and test that' ) raise Undecidable(msg) - elif isinstance(node, ParameterComparisonNode | DimensionComparisonNode | LookupComparisonNode): + elif isinstance(node, ParameterComparison | DimensionComparison | LookupComparison): if node.op not in ('==', '!=') and dtypes.get(subject.name) not in _ORDERED_DTYPES: msg = ( f'{subject} has dtype {dtypes.get(subject.name)!r} and is ordered with ' @@ -210,19 +210,19 @@ def _observe(node: TypedPredicateNode, subject: Subject, values: set[Any], dtype values.add(node.value) -def _subject_of(node: TypedPredicateNode) -> Subject: +def _subject_of(node: TypedPredicate) -> Subject: match node: - case ParameterDefinedNode(name=name) | ParameterComparisonNode(name=name): + case ParameterDefined(name=name) | ParameterComparison(name=name): return Subject('param', name) - case VariableDefinedNode(name=name): + case VariableDefined(name=name): return Subject('variable', name) - case DimensionComparisonNode(name=name): + case DimensionComparison(name=name): return Subject('dim', name) - case DimensionPositionNode(name=name, by=by): + case DimensionPosition(name=name, by=by): return Subject('rank', name, by) - case LookupDefinedNode(name=name) | LookupComparisonNode(name=name): + case LookupDefined(name=name) | LookupComparison(name=name): return Subject('lookup', name) - case LookupPairComparisonNode(name=name, other=other): + case LookupPairComparison(name=name, other=other): return Subject('lookup_pair', name, other) case _: assert_never(node) @@ -376,42 +376,42 @@ def _shown(subject: Subject, value: Cell) -> str: # --------------------------------------------------------------------------- -def _evaluate(node: WhereNode, cell: dict[Subject, Cell], grid: _Grid) -> bool: +def _evaluate(node: Predicate, cell: dict[Subject, Cell], grid: _Grid) -> bool: """Is *node* true in this cell?""" - if isinstance(node, TypedPredicateNode): + if isinstance(node, TypedPredicate): return _atom(node, cell, grid) match node: - case BooleanLiteralNode(value=value): + case BooleanLiteral(value=value): return value - case NotNode(operand=operand): + case Not(operand=operand): return not _evaluate(operand, cell, grid) - case AndNode(left=left, right=right): + case And(left=left, right=right): return _evaluate(left, cell, grid) and _evaluate(right, cell, grid) - case OrNode(left=left, right=right): + case Or(left=left, right=right): return _evaluate(left, cell, grid) or _evaluate(right, cell, grid) case _: assert_never(node) -def _atom(node: TypedPredicateNode, cell: dict[Subject, Cell], grid: _Grid) -> bool: +def _atom(node: TypedPredicate, cell: dict[Subject, Cell], grid: _Grid) -> bool: subject = grid.subjects[id(node)] value = cell[subject] match node: - case ParameterDefinedNode() | LookupDefinedNode(): + case ParameterDefined() | LookupDefined(): if isinstance(value, bool): return value return value not in (Special.NULL, Special.POS_INF, Special.NEG_INF) - case VariableDefinedNode(): + case VariableDefined(): return bool(value) - case LookupPairComparisonNode(op=op): + case LookupPairComparison(op=op): return bool(value) if op == '==' else not value - case DimensionPositionNode(op=op, position=position): + case DimensionPosition(op=op, position=position): return _compare(value, op, position) - case ParameterComparisonNode(op=op, value=literal) | LookupComparisonNode(op=op, value=literal): + case ParameterComparison(op=op, value=literal) | LookupComparison(op=op, value=literal): if value is Special.NULL: return False return _compare(value, op, literal) - case DimensionComparisonNode(op=op, value=literal): + case DimensionComparison(op=op, value=literal): return _compare(value, op, literal) case _: assert_never(node) diff --git a/src/math_spec/lowering.py b/src/math_spec/lowering.py index d1df7af4..2517cdec 100644 --- a/src/math_spec/lowering.py +++ b/src/math_spec/lowering.py @@ -140,15 +140,8 @@ def lower_program(expanded: _ExpandedSpec) -> program.Program: _Lowering(expanded, 'the objective').expr(resolved.objective), ) - dimensions = { - dname: program.DimensionDeclaration( - tuple( - program.LookupDeclaration(lname, lk.into) for lname, lk in expanded.lookups.items() if lk.over == dname - ), - ddef.dtype, - ) - for dname, ddef in expanded.dimensions.items() - } + dimensions = {dname: program.DimensionDeclaration(ddef.dtype) for dname, ddef in expanded.dimensions.items()} + lookups = {lname: program.LookupDeclaration(lk.over, lk.into) for lname, lk in expanded.lookups.items()} sos = { sname: program.SosDeclaration( sdef.variable, @@ -169,9 +162,10 @@ def lower_program(expanded: _ExpandedSpec) -> program.Program: constraints=constraints, objective=objective, dimensions=dimensions, + lookups=lookups, sos=sos, piecewise={name: declaration_of(ex) for name, ex in expanded.expanded_piecewise.items()}, - named_expressions=expressions, + expressions=expressions, ) @@ -187,7 +181,7 @@ class _Lowering: schema: _ExpandedSpec context: str - def expr(self, node: ArithmeticNode) -> program.ExpressionNode: + def expr(self, node: ArithmeticNode) -> program.Expression: """Rewrite one resolved core-AST expression as a program expression.""" if isinstance(node, NumberNode): return program.Constant(node.value) @@ -256,7 +250,7 @@ def _cases(self, node: CasesNode) -> program.Cases: regions.append(program.Region(when, self.expr(arm.value))) return program.Cases(tuple(regions)) - def sum(self, node: FunctionCallNode) -> program.ExpressionNode: + def sum(self, node: FunctionCallNode) -> program.Expression: """``sum(x)``, ``sum(x, over=d)`` or ``sum(x, by=lookup)``. Two program nodes under one surface verb: reducing a dim away and reducing it @@ -274,18 +268,18 @@ def sum(self, node: FunctionCallNode) -> program.ExpressionNode: assert isinstance(by_node, LookupNode), 'resolution refuses a by= that is not a lookup' return program.GroupSum(operand, over=by_node.dimension, coordinate=by_node.names, into=by_node.into) - def at(self, node: FunctionCallNode) -> program.ExpressionNode: + def at(self, node: FunctionCallNode) -> program.Expression: """``at(x, by=lookup)`` — the adjoint of :meth:`sum`'s ``by=`` form.""" by_node = node.kwargs['by'] assert isinstance(by_node, LookupNode), 'resolution refuses a by= that is not a lookup' - return program.At( + return program.Pullback( self.expr(node.args[0]), over=by_node.dimension, coordinate=by_node.names, into=by_node.into, ) - def sum_back(self, node: FunctionCallNode) -> program.ExpressionNode: + def sum_back(self, node: FunctionCallNode) -> program.Expression: """``sum_back(x, over=d, within=w)`` — a trailing window along one dimension. *within* is an integer literal of at least one, or a parameter naming a @@ -307,9 +301,9 @@ def sum_back(self, node: FunctionCallNode) -> program.ExpressionNode: else: assert isinstance(within_node, NumberNode), 'a within= that is neither is refused at load' width = int(within_node.value) - return program.Window(operand, over_node.name, width=width, wrap=wrap, partition=_partition_of(node)) + return program.WindowSum(operand, over_node.name, width=width, wrap=wrap, partition=_partition_of(node)) - def shift(self, node: FunctionCallNode) -> program.ExpressionNode: + def shift(self, node: FunctionCallNode) -> program.Expression: """``shift(x, over=d, offset=n)`` — the value at *t - offset* along one dim. What the vacated positions contribute is ``edge=``'s to say, and the @@ -337,7 +331,7 @@ def shift(self, node: FunctionCallNode) -> program.ExpressionNode: #: One lowering per name in the language's ``BUILTIN_NAMES``. -_CALLS: dict[str, Callable[[_Lowering, FunctionCallNode], program.ExpressionNode]] = { +_CALLS: dict[str, Callable[[_Lowering, FunctionCallNode], program.Expression]] = { 'sum': _Lowering.sum, 'at': _Lowering.at, 'sum_back': _Lowering.sum_back, @@ -359,7 +353,7 @@ def _partition_of(node: FunctionCallNode) -> str | None: return by_node.names[0] -def _bound_expression(value: float | str) -> program.ExpressionNode: +def _bound_expression(value: float | str) -> program.Expression: if isinstance(value, str): return program.Parameter(value) return program.Constant(value) diff --git a/src/math_spec/program.py b/src/math_spec/program.py index 5478a31a..5689b94f 100644 --- a/src/math_spec/program.py +++ b/src/math_spec/program.py @@ -22,7 +22,7 @@ from collections.abc import Mapping from dataclasses import dataclass, field, fields, replace from functools import cached_property -from typing import TYPE_CHECKING, Literal, NamedTuple, assert_never, get_args +from typing import TYPE_CHECKING, Literal, assert_never, get_args import math_spec.model as _model from math_spec._expression_parser import ComparisonOperator @@ -38,55 +38,55 @@ __all__ = [ 'QUADRATIC_POSITIONS', 'Add', - 'AndNode', - 'At', + 'And', 'AtLeastTwo', - 'BooleanLiteralNode', + 'BooleanLiteral', 'Cases', 'Check', - 'ConnectiveWhereNode', + 'Connective', 'Constant', 'ConstraintDeclaration', 'ConstraintSense', 'Contiguous', 'Curved', 'Derivation', - 'DimensionComparisonNode', + 'DimensionComparison', 'DimensionDeclaration', 'DimensionDtype', - 'DimensionPositionNode', + 'DimensionPosition', 'Divide', 'Dual', 'Expression', 'ExpressionDeclaration', - 'ExpressionNode', 'FanIn', 'FirstOf', 'Footprint', 'GroupSum', 'Increasing', 'LastOf', - 'LookupComparisonNode', + 'LookupComparison', 'LookupDeclaration', - 'LookupDefinedNode', - 'LookupPairComparisonNode', + 'LookupDefined', + 'LookupPairComparison', 'Mask', 'MaskOf', 'Multiply', 'Negate', - 'NotNode', + 'Not', 'ObjectiveDeclaration', 'ObjectiveSense', - 'OrNode', + 'Or', 'Parameter', - 'ParameterComparisonNode', + 'ParameterComparison', 'ParameterDeclaration', - 'ParameterDefinedNode', + 'ParameterDefined', 'ParameterDtype', 'PiecewiseDeclaration', 'Power', + 'Predicate', 'PredicateOperator', 'Program', + 'Pullback', 'QuadraticPosition', 'Reach', 'Region', @@ -94,14 +94,13 @@ 'SosDeclaration', 'Sum', 'Translate', - 'TypedPredicateNode', + 'TypedPredicate', 'Variable', 'VariableAbsence', 'VariableDeclaration', - 'VariableDefinedNode', + 'VariableDefined', 'VariableDomain', - 'WhereNode', - 'Window', + 'WindowSum', 'carries_variable', 'check_message', 'children', @@ -155,47 +154,28 @@ @dataclass(frozen=True) -class Expression: - """Base class for expressions over variables and parameters. - - The degree rules (``math_spec.degree``) hold on every tree the math reads - — :attr:`Program.expressions`, a bound, and a named expression that is - ``in_math`` — affine but where a :class:`QuadraticPosition` admits a - :class:`Multiply` of two variable-carrying operands. A - :class:`ExpressionDeclaration` the math never reads is held to none of - them. No node records which tree it stands in. - """ - - def __add__(self: ExpressionNode, other: ExpressionNode) -> ExpressionNode: - return Add(self, other) - - def __mul__(self: ExpressionNode, other: ExpressionNode) -> ExpressionNode: - return Multiply(self, other) - - -@dataclass(frozen=True) -class Constant(Expression): +class Constant: """A scalar constant.""" value: float @dataclass(frozen=True) -class Parameter(Expression): +class Parameter: """A parameter reference — contributes to the constant part.""" name: str @dataclass(frozen=True) -class Variable(Expression): +class Variable: """A variable reference — one term per existing variable row.""" name: str @dataclass(frozen=True) -class Dual(Expression): +class Dual: """A constraint's dual — its shadow price, read after the solve. Stands only under an :class:`ExpressionDeclaration` the math never reads: @@ -209,30 +189,30 @@ class Dual(Expression): @dataclass(frozen=True) -class Negate(Expression): - operand: ExpressionNode +class Negate: + operand: Expression @dataclass(frozen=True) -class Add(Expression): - left: ExpressionNode - right: ExpressionNode +class Add: + left: Expression + right: Expression @dataclass(frozen=True) -class Multiply(Expression): +class Multiply: """Product of two operands. Affine where at least one factor is variable-free; degree 2 where neither is, which ``math_spec.degree`` admits in a :data:`QuadraticPosition` alone. """ - left: ExpressionNode - right: ExpressionNode + left: Expression + right: Expression @dataclass(frozen=True) -class Power(Expression): +class Power: """``base ** exponent``, both variable-free wherever the math reads it. The language refuses a variable anywhere under it (``math_spec.degree``), @@ -240,28 +220,28 @@ class Power(Expression): coordinate like any other parameter arithmetic. """ - base: ExpressionNode - exponent: ExpressionNode + base: Expression + exponent: Expression @dataclass(frozen=True) -class Divide(Expression): +class Divide: """Quotient ``numerator / divisor``, the divisor variable-free wherever the math reads it (``math_spec.degree``).""" - numerator: ExpressionNode - divisor: ExpressionNode + numerator: Expression + divisor: Expression @dataclass(frozen=True) -class Sum(Expression): +class Sum: """Sum ``operand`` over the named dims, removing them from the result.""" - operand: ExpressionNode + operand: Expression over: tuple[str, ...] @dataclass(frozen=True) -class GroupSum(Expression): +class GroupSum: """Sum ``operand`` through coordinates declared on dim ``over``. ``coordinate`` names lookups carried by dim ``over`` whose values are @@ -271,14 +251,14 @@ class GroupSum(Expression): consumed in a single join. """ - operand: ExpressionNode + operand: Expression over: str coordinate: tuple[str, ...] into: tuple[str, ...] @dataclass(frozen=True) -class At(Expression): +class Pullback: """Read ``operand`` through a lookup — the adjoint of :class:`GroupSum`. Same mapping table, walked the other way: ``GroupSum`` consumes ``over`` @@ -286,14 +266,14 @@ class At(Expression): join fans out, many ``over`` labels sharing one ``into`` tuple. """ - operand: ExpressionNode + operand: Expression over: str coordinate: tuple[str, ...] into: tuple[str, ...] @dataclass(frozen=True) -class Translate(Expression): +class Translate: """Re-index along one dimension: the result at *t* is ``operand`` at *t - offset*. ``wrap`` is ``edge='wrap'`` in the file: periodic, and stated on every @@ -302,16 +282,16 @@ class Translate(Expression): and contribute it. Always ``None`` under ``wrap``. ``offset`` is an integer, or the name of an integer parameter that does - not depend on ``dimension`` and carries its sign in the values. + not depend on ``over`` and carries its sign in the values. - ``partition`` names a lookup over ``dimension``, and the translation then + ``partition`` names a lookup over ``over``, and the translation then happens inside each group it makes: 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 lookup sends nowhere reaches nothing. """ - operand: ExpressionNode - dimension: str + operand: Expression + over: str offset: int | str wrap: bool fill: float | None = None @@ -319,7 +299,7 @@ class Translate(Expression): @dataclass(frozen=True) -class Window(Expression): +class WindowSum: """Sum ``operand`` over a trailing window along one dimension. The result at *t* is the sum of the operand at every position from @@ -339,8 +319,8 @@ class Window(Expression): coordinate the lookup places nowhere reaches nothing — not even itself. """ - operand: ExpressionNode - dimension: str + operand: Expression + over: str width: int | str wrap: bool partition: str | None = None @@ -355,11 +335,11 @@ class Region: """ when: Mask - value: ExpressionNode + value: Expression @dataclass(frozen=True) -class Cases(Expression): +class Cases: """A value defined by region — exactly one region applies at each coordinate. The regions are disjoint and total, so a consumer adds them rather than @@ -370,13 +350,17 @@ class Cases(Expression): regions: tuple[Region, ...] -#: Every expression node, as one type. The set is *closed* — nothing registers -#: into it — so a consumer that walks it ends in ``assert_never`` and a node -#: added without a branch is a type error at the site that must grow one, -#: rather than a ``LanguageError`` raised at the first model that uses it. -#: ``Expression`` stays the base class the nodes inherit and the operators are -#: declared on; this is what a walk *takes*. -ExpressionNode = ( +#: Every expression node, as one type — what a walk takes. The set is +#: *closed*: nothing registers into it, so a consumer that walks it ends in +#: ``assert_never`` and a node added without a branch is a type error at the +#: site that must grow one, rather than a ``LanguageError`` raised at the first +#: model that uses it. The degree rules (``math_spec.degree``) hold on every +#: tree the math reads — :attr:`Program.roots`, a bound, and a named expression +#: that is ``in_math`` — affine but where a :class:`QuadraticPosition` admits a +#: :class:`Multiply` of two variable-carrying operands. A +#: :class:`ExpressionDeclaration` the math never reads is held to none of them. +#: No node records which tree it stands in. +Expression = ( Constant | Parameter | Variable @@ -388,14 +372,14 @@ class Cases(Expression): | Divide | Sum | GroupSum - | At + | Pullback | Translate - | Window + | WindowSum | Cases ) -def fan_in(expression: ExpressionNode) -> FanIn: +def fan_in(expression: Expression) -> FanIn: """How *expression*'s output rows relate to its input slots. For the absence rules, both classes other than ``'one-to-one'`` sum @@ -403,17 +387,17 @@ def fan_in(expression: ExpressionNode) -> FanIn: """ if isinstance(expression, (Sum, GroupSum)): return 'many-to-one' - if isinstance(expression, Window): + if isinstance(expression, WindowSum): return 'one-to-many' if isinstance( expression, - (Constant, Parameter, Variable, Dual, Negate, Add, Multiply, Power, Divide, At, Translate, Cases), + (Constant, Parameter, Variable, Dual, Negate, Add, Multiply, Power, Divide, Pullback, Translate, Cases), ): return 'one-to-one' assert_never(expression) -def children(expression: ExpressionNode) -> tuple[ExpressionNode, ...]: +def children(expression: Expression) -> tuple[Expression, ...]: """The sub-expressions of *expression* — what every walk recurses through.""" if isinstance(expression, Negate): return (expression.operand,) @@ -423,7 +407,7 @@ def children(expression: ExpressionNode) -> tuple[ExpressionNode, ...]: return (expression.numerator, expression.divisor) if isinstance(expression, Power): return (expression.base, expression.exponent) - if isinstance(expression, (Sum, GroupSum, At, Translate, Window)): + if isinstance(expression, (Sum, GroupSum, Pullback, Translate, WindowSum)): return (expression.operand,) if isinstance(expression, Cases): return tuple(region.value for region in expression.regions) @@ -437,41 +421,30 @@ def children(expression: ExpressionNode) -> tuple[ExpressionNode, ...]: # -------------------------------------------------------------------------- -class LookupDeclaration(NamedTuple): - """One declared lookup over a dimension. +@dataclass(frozen=True) +class LookupDeclaration: + """One declared map out of dimension ``over`` into dimension ``into``. - Its values are labels of ``target``, checked for containment once the dim + Its values are labels of ``into``, checked for containment once the dim tables exist — which keeps a mistyped label from silently dropping its terms in the join that places them — and it is what ``sum(by=)`` lands terms on. """ - name: str - target: str + over: str + into: str @dataclass(frozen=True) class DimensionDeclaration: - """A dimension and the lookups its labels carry.""" + """A dimension, as the file declares it.""" - lookups: tuple[LookupDeclaration, ...] = () #: What the labels are, as the file declares them. A dimension is read from #: whatever table carries it, so the declared type is what that column is #: checked against — the same claim ``ParameterDeclaration.dtype`` makes #: about a value column, one axis over. dtype: DimensionDtype = 'str' - @property - def targets(self) -> dict[str, str]: - """Each map over the dimension, to the dimension its values are labels of. - - The question every consumer of a ``by=`` asks, and asked here so it has - one answer: an operator grouping through a lookup names the target as - the dim it lands on, and a partition array is named for it so an amount - declared over the group's own dim can be read through it. - """ - return {lk.name: lk.target for lk in self.lookups} - @dataclass(frozen=True) class MaskOf: @@ -635,8 +608,8 @@ class ParameterDeclaration: class VariableDeclaration: dims: tuple[str, ...] where: Mask | None = None - lower: ExpressionNode = field(default_factory=lambda: Constant(float('-inf'))) - upper: ExpressionNode = field(default_factory=lambda: Constant(float('inf'))) + lower: Expression = field(default_factory=lambda: Constant(float('-inf'))) + upper: Expression = field(default_factory=lambda: Constant(float('inf'))) domain: VariableDomain = 'continuous' absence: VariableAbsence = 'undefined' @@ -651,9 +624,9 @@ class ConstraintDeclaration: """ dims: tuple[str, ...] - lhs: ExpressionNode + lhs: Expression sense: ConstraintSense - rhs: ExpressionNode + rhs: Expression where: Mask | None = None @@ -665,7 +638,7 @@ class SosDeclaration: columns a consumer already has and says what may be nonzero among them. Which dims those are is the variable's own ``dims`` and is read from it: a copy here would be a second home for a fact - (:meth:`Program.variable`). + (:attr:`Program.variables`). ``big_m`` caps the linking coefficient a consumer without the concept reformulates with, and is ``None`` where the variable's own upper bound is @@ -683,7 +656,7 @@ class ObjectiveDeclaration: """Objective — scalar, every reduction in it one the file wrote.""" sense: ObjectiveSense - expression: ExpressionNode + expression: Expression @dataclass(frozen=True) @@ -692,13 +665,13 @@ class ExpressionDeclaration: ``in_math`` where the objective or a constraint inlines it, directly or through another entry or a macro; its body then stands inside - :attr:`Program.expressions` and is held to the degree rules where it is + :attr:`Program.roots` and is held to the degree rules where it is read. Otherwise nothing a solver sees contains it: it is a reported quantity, its body held to no degree, the one place a :class:`Dual` may stand. A bound and a ``where`` name no entry, so neither decides this. """ - expression: ExpressionNode + expression: Expression in_math: bool @@ -714,21 +687,13 @@ class Footprint: stands in; empty is affine throughout. domains: Every domain declared. sos_types: The order of each special-ordered set declared. - shapes: Every expression node kind that appears. + nodes: Every expression node kind that appears. """ quadratic: frozenset[QuadraticPosition] domains: frozenset[VariableDomain] sos_types: frozenset[Literal[1, 2]] - shapes: frozenset[type[ExpressionNode]] - - -def _declared[Declaration](items: Mapping[str, Declaration], name: str, kind: str) -> Declaration: - """The declaration called *name*, or a ``KeyError`` naming the near miss.""" - try: - return items[name] - except KeyError: - raise KeyError(f"unknown {kind} '{name}'. " + did_you_mean(name, list(items))) from None + nodes: frozenset[type[Expression]] @dataclass(frozen=True) @@ -847,6 +812,7 @@ class Program: #: whose answer is whether the constraints can be met at all. objective: ObjectiveDeclaration | None dimensions: Mapping[str, DimensionDeclaration] = Sealed({}) + lookups: Mapping[str, LookupDeclaration] = Sealed({}) sos: Mapping[str, SosDeclaration] = Sealed({}) #: Each ``piecewise:`` block the file wrote, as facts — see #: :class:`PiecewiseDeclaration`. @@ -856,7 +822,7 @@ class Program: #: it is read — but all are lowered with the program, so a file whose #: named expression is outside the language is refused by every verb that #: reads the file rather than only by the one that reads the expression. - named_expressions: Mapping[str, ExpressionDeclaration] = Sealed({}) + expressions: Mapping[str, ExpressionDeclaration] = Sealed({}) def __post_init__(self) -> None: """Seal every group, so a program handed out cannot be written to.""" @@ -865,16 +831,16 @@ def __post_init__(self) -> None: if isinstance(group, Mapping): object.__setattr__(self, f.name, Sealed(group)) - def _by_position(self) -> Iterator[tuple[QuadraticPosition, tuple[ExpressionNode, ...]]]: + def _by_position(self) -> Iterator[tuple[QuadraticPosition, tuple[Expression, ...]]]: """The row-building expressions, grouped by the position they stand in.""" yield 'objective', (self.objective.expression,) if self.objective is not None else () yield 'constraint', tuple(side for c in self.constraints.values() for side in (c.lhs, c.rhs)) @property - def expressions(self) -> tuple[ExpressionNode, ...]: - """Every expression a row is built from — the objective and both sides of each constraint. + def roots(self) -> tuple[Expression, ...]: + """Every tree a row is built from — the objective and both sides of each constraint. - A :attr:`named_expressions` entry builds no row and is not among them. + An :attr:`expressions` entry builds no row and is not among them. """ return tuple(e for _, group in self._by_position() for e in group) @@ -887,23 +853,9 @@ def footprint(self) -> Footprint: ), domains=frozenset(v.domain for v in self.variables.values()), sos_types=frozenset(s.sos_type for s in self.sos.values()), - shapes=frozenset(type(node) for node in walk(*self.expressions)), + nodes=frozenset(type(node) for node in walk(*self.roots)), ) - def dimension(self, name: str) -> DimensionDeclaration: - return _declared(self.dimensions, name, 'dimension') - - @property - def lookups(self) -> tuple[tuple[str, LookupDeclaration], ...]: - """Every lookup in the program, with the dimension it is over.""" - return tuple((dimension, lk) for dimension, d in self.dimensions.items() for lk in d.lookups) - - def parameter(self, name: str) -> ParameterDeclaration: - return _declared(self.parameters, name, 'parameter') - - def variable(self, name: str) -> VariableDeclaration: - return _declared(self.variables, name, 'variable') - @cached_property def separability(self) -> Mapping[str, Separability]: """Every axis, to what building it a window at a time asks and what it would break. @@ -934,7 +886,7 @@ def separability(self) -> Mapping[str, Separability]: # -------------------------------------------------------------------------- -def walk(*expressions: ExpressionNode) -> Iterator[ExpressionNode]: +def walk(*expressions: Expression) -> Iterator[Expression]: """Every node under *expressions*, each expression itself included, parents first. The traversal every *question* about a program is a filter of — which names @@ -949,7 +901,7 @@ def walk(*expressions: ExpressionNode) -> Iterator[ExpressionNode]: yield from walk(*children(expression)) -def is_quadratic(expression: ExpressionNode) -> bool: +def is_quadratic(expression: Expression) -> bool: """Whether *expression* contains a product of two variable-carrying operands. A structural question over the program, and unrelated consumers ask it — @@ -968,22 +920,22 @@ def is_quadratic(expression: ExpressionNode) -> bool: ) -def carries_variable(expression: ExpressionNode) -> bool: +def carries_variable(expression: Expression) -> bool: """Whether a variable appears anywhere under *expression*.""" return any(isinstance(node, Variable) for node in walk(expression)) -def parameters_of(*expressions: ExpressionNode) -> frozenset[str]: +def parameters_of(*expressions: Expression) -> frozenset[str]: """Every parameter named anywhere under *expressions*.""" return frozenset(node.name for node in walk(*expressions) if isinstance(node, Parameter)) -def variables_of(*expressions: ExpressionNode) -> frozenset[str]: +def variables_of(*expressions: Expression) -> frozenset[str]: """Every variable named anywhere under *expressions*.""" return frozenset(node.name for node in walk(*expressions) if isinstance(node, Variable)) -def quotients(*expressions: ExpressionNode) -> tuple[Divide, ...]: +def quotients(*expressions: Expression) -> tuple[Divide, ...]: """Every division under *expressions*, each kept whole. The divisor and the numerator answer different questions and one consumer @@ -994,7 +946,7 @@ def quotients(*expressions: ExpressionNode) -> tuple[Divide, ...]: return tuple(node for node in walk(*expressions) if isinstance(node, Divide)) -def divisor_parameters(*expressions: ExpressionNode) -> frozenset[str]: +def divisor_parameters(*expressions: Expression) -> frozenset[str]: """Every parameter named anywhere in a divisor under *expressions*.""" return frozenset().union(*(parameters_of(q.divisor) for q in quotients(*expressions))) @@ -1008,12 +960,12 @@ def divisor_parameters(*expressions: ExpressionNode) -> frozenset[str]: @dataclass(frozen=True) -class BooleanLiteralNode: +class BooleanLiteral: value: bool @dataclass(frozen=True) -class ParameterDefinedNode: +class ParameterDefined: """True wherever the named parameter is non-null and finite. ``dims`` is the parameter's own, copied off the declaration during @@ -1026,7 +978,7 @@ class ParameterDefinedNode: @dataclass(frozen=True) -class VariableDefinedNode: +class VariableDefined: """True at the coordinates where the named variable exists.""" name: str @@ -1034,7 +986,7 @@ class VariableDefinedNode: @dataclass(frozen=True) -class ParameterComparisonNode: +class ParameterComparison: """Compare a parameter against a literal, element-wise.""" name: str @@ -1044,7 +996,7 @@ class ParameterComparisonNode: @dataclass(frozen=True) -class DimensionComparisonNode: +class DimensionComparison: """Compare a dimension's own coordinates against a literal.""" name: str @@ -1053,7 +1005,7 @@ class DimensionComparisonNode: @dataclass(frozen=True) -class DimensionPositionNode: +class DimensionPosition: """Compare where a row sits along a dimension against a position — ``position(snapshot) == 0``. Both sides are integers, negative counting from the end. With ``by`` the @@ -1067,7 +1019,7 @@ class DimensionPositionNode: @dataclass(frozen=True) -class LookupComparisonNode: +class LookupComparison: """Compare a lookup's values against a literal — ``period_of == 2030``. ``over`` is the dimension the lookup maps out of. @@ -1080,7 +1032,7 @@ class LookupComparisonNode: @dataclass(frozen=True) -class LookupPairComparisonNode: +class LookupPairComparison: """Compare two lookups over one dimension — ``from != to``, row by row on that dimension's table.""" name: str @@ -1090,7 +1042,7 @@ class LookupPairComparisonNode: @dataclass(frozen=True) -class LookupDefinedNode: +class LookupDefined: """True where the named lookup has a value — a null says the label belongs to no group.""" name: str @@ -1098,77 +1050,77 @@ class LookupDefinedNode: @dataclass(frozen=True) -class NotNode: - operand: WhereNode +class Not: + operand: Predicate @dataclass(frozen=True) -class AndNode: - left: WhereNode - right: WhereNode +class And: + left: Predicate + right: Predicate @dataclass(frozen=True) -class OrNode: - left: WhereNode - right: WhereNode +class Or: + left: Predicate + right: Predicate #: Every resolved predicate node — what a lowered mask's ``root`` is built of. #: The parser's ``Unresolved*`` nodes are not members: they live with the #: grammar in :mod:`math_spec._where_parser`, and resolution rewrites them away #: before anything here is asked. -WhereNode = ( - BooleanLiteralNode - | DimensionPositionNode - | ParameterDefinedNode - | VariableDefinedNode - | ParameterComparisonNode - | DimensionComparisonNode - | LookupComparisonNode - | LookupPairComparisonNode - | LookupDefinedNode - | NotNode - | AndNode - | OrNode +Predicate = ( + BooleanLiteral + | DimensionPosition + | ParameterDefined + | VariableDefined + | ParameterComparison + | DimensionComparison + | LookupComparison + | LookupPairComparison + | LookupDefined + | Not + | And + | Or ) #: Every predicate resolution has typed: it names a declaration and the kind is #: settled. Resolution passes these straight through, having nothing left to #: decide about them. -TypedPredicateNode = ( - ParameterComparisonNode - | ParameterDefinedNode - | VariableDefinedNode - | DimensionComparisonNode - | DimensionPositionNode - | LookupComparisonNode - | LookupPairComparisonNode - | LookupDefinedNode +TypedPredicate = ( + ParameterComparison + | ParameterDefined + | VariableDefined + | DimensionComparison + | DimensionPosition + | LookupComparison + | LookupPairComparison + | LookupDefined ) #: The boolean connectives — the only where nodes carrying other where nodes, #: and so the only place a walk over a predicate recurses. The grammar builds #: these classes directly, over leaves still unresolved, so a pre-resolution #: tree shares them — the transient impurity resolution normalizes away. -ConnectiveWhereNode = NotNode | AndNode | OrNode +Connective = Not | And | Or -def where_children(where: WhereNode) -> tuple[WhereNode, ...]: +def where_children(where: Predicate) -> tuple[Predicate, ...]: """The predicates under *where* — a connective's operands, and nothing under a leaf. What every walk over a predicate recurses through, as :func:`children` is for an expression. A leaf has nothing under it whether or not it is resolved, so the grammar measures its own output with this too. """ - if isinstance(where, NotNode): + if isinstance(where, Not): return (where.operand,) - if isinstance(where, (AndNode, OrNode)): + if isinstance(where, (And, Or)): return (where.left, where.right) return () -def _atoms(where: WhereNode) -> Iterator[TypedPredicateNode]: +def _atoms(where: Predicate) -> Iterator[TypedPredicate]: """Every node in *where* that reads a declaration, connectives removed. A boolean literal yields nothing. @@ -1176,9 +1128,9 @@ def _atoms(where: WhereNode) -> Iterator[TypedPredicateNode]: Raises: AssertionError: An unresolved node reached the walk. """ - if isinstance(where, TypedPredicateNode): + if isinstance(where, TypedPredicate): yield where - elif isinstance(where, BooleanLiteralNode | ConnectiveWhereNode): + elif isinstance(where, BooleanLiteral | Connective): for child in where_children(where): yield from _atoms(child) else: @@ -1186,7 +1138,7 @@ def _atoms(where: WhereNode) -> Iterator[TypedPredicateNode]: raise AssertionError(msg) -def _atom_dims(atom: TypedPredicateNode) -> frozenset[str]: +def _atom_dims(atom: TypedPredicate) -> frozenset[str]: """One leaf's dims — the rule :attr:`Mask.dims` is the union of. A parameter or variable leaf carries its own dims off the declaration; a @@ -1198,17 +1150,17 @@ def _atom_dims(atom: TypedPredicateNode) -> frozenset[str]: a branch, rather than a wrong dim set at the first model to use it. """ match atom: - case ParameterComparisonNode() | ParameterDefinedNode() | VariableDefinedNode(): + case ParameterComparison() | ParameterDefined() | VariableDefined(): return frozenset(atom.dims) - case DimensionComparisonNode() | DimensionPositionNode(): + case DimensionComparison() | DimensionPosition(): return frozenset({atom.name}) - case LookupComparisonNode() | LookupPairComparisonNode() | LookupDefinedNode(): + case LookupComparison() | LookupPairComparison() | LookupDefined(): return frozenset({atom.over}) case _: assert_never(atom) -def _atom_names(atom: TypedPredicateNode) -> frozenset[str]: +def _atom_names(atom: TypedPredicate) -> frozenset[str]: """One leaf's declarations, its dimension apart — the rule :attr:`Mask.names_read` is the union of. A comparison on a dimension names no declaration — a coordinate is not @@ -1218,23 +1170,17 @@ def _atom_names(atom: TypedPredicateNode) -> frozenset[str]: than a name silently dropped at the first model to use it. """ match atom: - case ( - ParameterComparisonNode() - | ParameterDefinedNode() - | VariableDefinedNode() - | LookupComparisonNode() - | LookupDefinedNode() - ): + case ParameterComparison() | ParameterDefined() | VariableDefined() | LookupComparison() | LookupDefined(): return frozenset({atom.name}) - case LookupPairComparisonNode(): + case LookupPairComparison(): return frozenset({atom.name, atom.other}) - case DimensionComparisonNode() | DimensionPositionNode(): + case DimensionComparison() | DimensionPosition(): return frozenset() case _: assert_never(atom) -def _conjuncts(where: WhereNode) -> tuple[WhereNode, ...]: +def _conjuncts(where: Predicate) -> tuple[Predicate, ...]: """The flatten rule behind :attr:`Mask.conjuncts` — the one home of the split. ``a AND b AND c`` gives three, and a predicate that is not an ``AND`` gives @@ -1243,12 +1189,12 @@ def _conjuncts(where: WhereNode) -> tuple[WhereNode, ...]: ``NOT (a AND b)`` the single ``NOT`` — neither an ``OR`` nor a ``NOT`` is a claim the predicate makes on its own, so neither is split. """ - if isinstance(where, AndNode): + if isinstance(where, And): return _conjuncts(where.left) + _conjuncts(where.right) return (where,) -def _fold(node: WhereNode) -> WhereNode: +def _fold(node: Predicate) -> Predicate: """*node* with every connective a literal or a double negation decides evaluated away. ``X AND True`` is ``X``, ``X OR True`` is every row, ``X AND False`` is @@ -1257,27 +1203,27 @@ def _fold(node: WhereNode) -> WhereNode: the invariant :class:`Mask` applies at construction, so it holds wherever a mask is built. """ - if isinstance(node, NotNode): + if isinstance(node, Not): operand = _fold(node.operand) - if isinstance(operand, BooleanLiteralNode): - return BooleanLiteralNode(not operand.value) - if isinstance(operand, NotNode): + if isinstance(operand, BooleanLiteral): + return BooleanLiteral(not operand.value) + if isinstance(operand, Not): return operand.operand - return NotNode(operand) - if isinstance(node, AndNode): + return Not(operand) + if isinstance(node, And): left, right = _fold(node.left), _fold(node.right) - if isinstance(left, BooleanLiteralNode): + if isinstance(left, BooleanLiteral): return right if left.value else left - if isinstance(right, BooleanLiteralNode): + if isinstance(right, BooleanLiteral): return left if right.value else right - return AndNode(left, right) - if isinstance(node, OrNode): + return And(left, right) + if isinstance(node, Or): left, right = _fold(node.left), _fold(node.right) - if isinstance(left, BooleanLiteralNode): + if isinstance(left, BooleanLiteral): return left if left.value else right - if isinstance(right, BooleanLiteralNode): + if isinstance(right, BooleanLiteral): return right if right.value else left - return OrNode(left, right) + return Or(left, right) return node @@ -1294,14 +1240,14 @@ class Mask: root: The resolved predicate the mask restricts rows by, folded. """ - root: WhereNode + root: Predicate def __post_init__(self) -> None: object.__setattr__(self, 'root', _fold(self.root)) _ = self.atoms # the walk is the refusal, and runs after the fold @cached_property - def atoms(self) -> tuple[TypedPredicateNode, ...]: + def atoms(self) -> tuple[TypedPredicate, ...]: """The mask's leaves, connectives removed — the one walk the other questions read. Held rather than re-walked: construction takes this walk anyway, to @@ -1310,7 +1256,7 @@ def atoms(self) -> tuple[TypedPredicateNode, ...]: return tuple(_atoms(self.root)) @property - def conjuncts(self) -> tuple[WhereNode, ...]: + def conjuncts(self) -> tuple[Predicate, ...]: """The predicates the mask joins with ``AND`` — its ``AND`` spine flattened, stopping at an ``OR`` or a ``NOT``.""" return _conjuncts(self.root) @@ -1331,12 +1277,12 @@ def dims(self) -> frozenset[str]: def __invert__(self) -> Mask: """The mask admitting exactly the rows this one refuses — construction folds a double negation or a literal flip.""" - return Mask(NotNode(self.root)) + return Mask(Not(self.root)) def __and__(self, other: Mask) -> Mask: """Both masks at once — construction absorbs a literal side rather than burying it.""" - return Mask(AndNode(self.root, other.root)) + return Mask(And(self.root, other.root)) def __or__(self, other: Mask) -> Mask: """Either mask — construction absorbs a literal side rather than burying it.""" - return Mask(OrNode(self.root, other.root)) + return Mask(Or(self.root, other.root)) diff --git a/src/math_spec/resolution.py b/src/math_spec/resolution.py index 4af355b0..33659cea 100644 --- a/src/math_spec/resolution.py +++ b/src/math_spec/resolution.py @@ -61,21 +61,21 @@ unknown_operator_message, ) from math_spec.program import ( - AndNode, - BooleanLiteralNode, - DimensionComparisonNode, - DimensionPositionNode, - LookupComparisonNode, - LookupDefinedNode, - LookupPairComparisonNode, + And, + BooleanLiteral, + DimensionComparison, + DimensionPosition, + LookupComparison, + LookupDefined, + LookupPairComparison, Mask, - NotNode, - OrNode, - ParameterComparisonNode, - ParameterDefinedNode, - TypedPredicateNode, - VariableDefinedNode, - WhereNode, + Not, + Or, + ParameterComparison, + ParameterDefined, + Predicate, + TypedPredicate, + VariableDefined, ) if TYPE_CHECKING: @@ -277,7 +277,7 @@ def where_of(text: str | None, ns: Namespace, context: str, self_variable: str | ``None`` for no mask, however the file spelled it: a mask that admits every row is dropped, and one that admits none arrives as a mask over - ``BooleanLiteralNode(False)``. + ``BooleanLiteral(False)``. Raises: LanguageError: Listing every problem the predicate has. @@ -296,9 +296,9 @@ def names_in(value: ArithmeticNode) -> tuple[str, ...]: return value.names if isinstance(value, NameListNode) else () -def mask_of(node: WhereNode | None) -> Mask | None: +def mask_of(node: Predicate | None) -> Mask | None: """The mask a declaration carries for a resolved where: ``None`` where there is none, or where every row passes.""" - if node is None or (isinstance(node, BooleanLiteralNode) and node.value): + if node is None or (isinstance(node, BooleanLiteral) and node.value): return None return Mask(node) @@ -327,22 +327,22 @@ def resolve_expression( def resolve_where( - node: WhereNode | UnresolvedWhereNode, + node: Predicate | UnresolvedWhereNode, ns: Namespace, context: str, errors: list[str], self_variable: str | None = None, -) -> WhereNode | None: +) -> Predicate | None: """Rewrite a parsed where AST into typed predicates, folded as :class:`~math_spec.program.Mask` folds. Returns: The typed tree — a mask admitting every row or none comes back as the - one ``BooleanLiteralNode`` — or ``None`` once anything failed, with the + one ``BooleanLiteral`` — or ``None`` once anything failed, with the problems appended to *errors*. """ before = len(errors) resolved = _Resolver(ns, context, errors, self_variable).where(node) - return None if len(errors) > before else Mask(cast('WhereNode', resolved)).root + return None if len(errors) > before else Mask(cast('Predicate', resolved)).root def resolve_where_text( @@ -351,7 +351,7 @@ def resolve_where_text( context: str, errors: list[str], self_variable: str | None = None, -) -> WhereNode | None: +) -> Predicate | None: """Parse and resolve one where string as :func:`resolve_where` does, a parse failure appended to *errors*. Returns: @@ -624,9 +624,9 @@ def _not_a_lookup(self, name: str, operator: str, key: str) -> str | None: # -- where strings ----------------------------------------------------- - def where(self, node: WhereNode | UnresolvedWhereNode) -> WhereNode | UnresolvedWhereNode: + def where(self, node: Predicate | UnresolvedWhereNode) -> Predicate | UnresolvedWhereNode: """One predicate node typed, or returned unresolved with its refusal appended.""" - if isinstance(node, BooleanLiteralNode | TypedPredicateNode): + if isinstance(node, BooleanLiteral | TypedPredicate): return node if isinstance(node, UnresolvedNameNode): return self._where_name(node) @@ -634,19 +634,19 @@ def where(self, node: WhereNode | UnresolvedWhereNode) -> WhereNode | Unresolved return self._position(node) if isinstance(node, UnresolvedComparisonNode): return self._comparison(node) - if isinstance(node, NotNode): - return NotNode(self._child(node.operand)) - if isinstance(node, AndNode): - return AndNode(self._child(node.left), self._child(node.right)) - if isinstance(node, OrNode): - return OrNode(self._child(node.left), self._child(node.right)) + if isinstance(node, Not): + return Not(self._child(node.operand)) + if isinstance(node, And): + return And(self._child(node.left), self._child(node.right)) + if isinstance(node, Or): + return Or(self._child(node.left), self._child(node.right)) assert_never(node) - def _child(self, node: WhereNode | UnresolvedWhereNode) -> WhereNode: + def _child(self, node: Predicate | UnresolvedWhereNode) -> Predicate: """A connective's child, typed as resolved: an unresolved one survives only with its refusal appended.""" - return cast('WhereNode', self.where(node)) + return cast('Predicate', self.where(node)) - def _where_name(self, node: UnresolvedNameNode) -> WhereNode | UnresolvedWhereNode: + def _where_name(self, node: UnresolvedNameNode) -> Predicate | UnresolvedWhereNode: """A bare name: a parameter's or lookup's definedness, or a variable's existence.""" ns, context = self.ns, self.context kind = ns.kind(node.name) @@ -655,7 +655,7 @@ def _where_name(self, node: UnresolvedNameNode) -> WhereNode | UnresolvedWhereNo return node match kind: case 'parameter': - return ParameterDefinedNode(node.name, ns.leaf_dims[node.name]) + return ParameterDefined(node.name, ns.leaf_dims[node.name]) case 'dimension': self.errors.append( f"{context}: '{node.name}' is a dimension, and a bare dimension " @@ -663,7 +663,7 @@ def _where_name(self, node: UnresolvedNameNode) -> WhereNode | UnresolvedWhereNo f'Remove it, or compare it: where: "{node.name} > 0".' ) case 'lookup': - return LookupDefinedNode(node.name, ns.over_of(node.name)) + return LookupDefined(node.name, ns.over_of(node.name)) case 'variable': if node.name == self.self_variable: self.errors.append( @@ -672,10 +672,10 @@ def _where_name(self, node: UnresolvedNameNode) -> WhereNode | UnresolvedWhereNo f'exists. Test a parameter, or another variable declared before it.' ) else: - return VariableDefinedNode(node.name, ns.leaf_dims[node.name]) + return VariableDefined(node.name, ns.leaf_dims[node.name]) return node - def _position(self, node: UnresolvedPositionNode) -> DimensionPositionNode | UnresolvedPositionNode: + def _position(self, node: UnresolvedPositionNode) -> DimensionPosition | UnresolvedPositionNode: """``position(dim[, by=lookup]) i``: the name a dimension, ``by=`` a lookup over it.""" ns, context = self.ns, self.context if node.dimension not in ns.dimensions: @@ -686,7 +686,7 @@ def _position(self, node: UnresolvedPositionNode) -> DimensionPositionNode | Unr ) return node if node.by is None: - return DimensionPositionNode(node.dimension, node.op, node.position, node.by) + return DimensionPosition(node.dimension, node.op, node.position, node.by) call = f'position({node.dimension}, by={node.by})' if ns.kind(node.by) != 'lookup': self.errors.append( @@ -703,9 +703,9 @@ def _position(self, node: UnresolvedPositionNode) -> DimensionPositionNode | Unr f"position within a group to name — group by a lookup over '{node.dimension}'." ) return node - return DimensionPositionNode(node.dimension, node.op, node.position, node.by) + return DimensionPosition(node.dimension, node.op, node.position, node.by) - def _comparison(self, node: UnresolvedComparisonNode) -> WhereNode | UnresolvedWhereNode: + def _comparison(self, node: UnresolvedComparisonNode) -> Predicate | UnresolvedWhereNode: """``name literal``, or the one structural form ``lookup lookup``.""" ns, context = self.ns, self.context value = node.value @@ -714,7 +714,7 @@ def _comparison(self, node: UnresolvedComparisonNode) -> WhereNode | UnresolvedW if (refusal := _lookup_pair_error(context, node, value, ns)) is not None: self.errors.append(refusal) return node - return LookupPairComparisonNode(node.name, value, ns.over_of(node.name), node.op) + return LookupPairComparison(node.name, value, ns.over_of(node.name), node.op) self.errors.append(_declared_rhs_error(context, node, value, rhs_kind)) return node @@ -731,11 +731,11 @@ def _comparison(self, node: UnresolvedComparisonNode) -> WhereNode | UnresolvedW match kind: case 'parameter': assert not isinstance(value, datetime.date) - return ParameterComparisonNode(node.name, node.op, value, ns.leaf_dims[node.name]) + return ParameterComparison(node.name, node.op, value, ns.leaf_dims[node.name]) case 'dimension': - return DimensionComparisonNode(node.name, node.op, value) + return DimensionComparison(node.name, node.op, value) case 'lookup': - return LookupComparisonNode(node.name, ns.over_of(node.name), node.op, value) + return LookupComparison(node.name, ns.over_of(node.name), node.op, value) case 'variable': self.errors.append( f"{context}: where references variable '{node.name}'. A where " diff --git a/src/math_spec/separability.py b/src/math_spec/separability.py index 64a2ac12..96a5faf6 100644 --- a/src/math_spec/separability.py +++ b/src/math_spec/separability.py @@ -9,26 +9,26 @@ from typing import TYPE_CHECKING, Literal from math_spec.program import ( - At, Cases, - DimensionPositionNode, + DimensionPosition, GroupSum, Mask, + Pullback, Reach, Separability, Sum, Translate, - Window, + WindowSum, walk, ) if TYPE_CHECKING: from collections.abc import Iterator - from math_spec.program import ExpressionNode, Program + from math_spec.program import Expression, Program -def _built_blocks(program: Program) -> Iterator[tuple[str, tuple[ExpressionNode, ...], Mask | None, bool]]: +def _built_blocks(program: Program) -> Iterator[tuple[str, tuple[Expression, ...], Mask | None, bool]]: """Every block that builds rows, labelled as the lowering's own messages label it. A named expression is not one: it is inlined where it is referenced, so @@ -90,12 +90,12 @@ def waits_on(dimension: str, label: str, name: str, kind: Literal['offset', 'par label, f'groups {node.over} into {", ".join(node.into)} — window that dimension instead, or cut only at the group edges', ) - elif isinstance(node, At): + elif isinstance(node, Pullback): for dimension in node.into: for lookup in node.coordinate: waits_on(dimension, label, lookup, 'coordinate') - elif isinstance(node, (Translate, Window)): - dimension = node.dimension + elif isinstance(node, (Translate, WindowSum)): + dimension = node.over if node.wrap: report( 'coupled', @@ -107,7 +107,7 @@ def waits_on(dimension: str, label: str, name: str, kind: Literal['offset', 'par continue if node.partition is not None: waits_on(dimension, label, node.partition, 'partition') - if isinstance(node, Window): + if isinstance(node, WindowSum): continue if isinstance(node.offset, str): waits_on(dimension, label, node.offset, 'offset') @@ -115,7 +115,7 @@ def waits_on(dimension: str, label: str, name: str, kind: Literal['offset', 'par 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, DimensionPositionNode): + if isinstance(atom, DimensionPosition): report('restarts', atom.name, label, f'counts a position along {atom.name}') for name, block in program.sos.items(): diff --git a/src/math_spec/typesetting/walk.py b/src/math_spec/typesetting/walk.py index 4569052f..ad90c03b 100644 --- a/src/math_spec/typesetting/walk.py +++ b/src/math_spec/typesetting/walk.py @@ -35,21 +35,21 @@ ) from math_spec.dimensions import dims_of from math_spec.program import ( - AndNode, - BooleanLiteralNode, - DimensionComparisonNode, - DimensionPositionNode, - LookupComparisonNode, - LookupDefinedNode, - LookupPairComparisonNode, + And, + BooleanLiteral, + DimensionComparison, + DimensionPosition, + LookupComparison, + LookupDefined, + LookupPairComparison, Mask, - NotNode, - OrNode, - ParameterComparisonNode, - ParameterDefinedNode, + Not, + Or, + ParameterComparison, + ParameterDefined, + Predicate, PredicateOperator, - VariableDefinedNode, - WhereNode, + VariableDefined, ) from math_spec.typesetting.format import Entry, Glossary, Line, OperatorName @@ -497,33 +497,33 @@ def _reduction_body(self, node: ArithmeticNode, ctx: _Context) -> str: # -- where strings ----------------------------------------------------- - def _predicate(self, node: WhereNode, ctx: _Context, *, need: int = 0) -> str: + def _predicate(self, node: Predicate, ctx: _Context, *, need: int = 0) -> str: text, precedence = self._where(node, ctx) return self.format.parenthesise(text) if precedence < need else text - def _where(self, node: WhereNode, ctx: _Context) -> tuple[str, int]: + def _where(self, node: Predicate, ctx: _Context) -> tuple[str, int]: comparison = _WHERE_PRECEDENCE['comparison'] - if isinstance(node, BooleanLiteralNode): + if isinstance(node, BooleanLiteral): assert not node.value, 'an always-true mask is folded away or refused before anything prints it' return self._op('false'), _ATOM - if isinstance(node, ParameterDefinedNode): + if isinstance(node, ParameterDefined): indexed = ctx.indexed(self.symbols.name[node.name], list(node.dims)) if self.schema.parameters[node.name].dtype == 'bool': return indexed, _ATOM return f'{indexed} {self.format.prose(" is defined")}', comparison - if isinstance(node, VariableDefinedNode): + if isinstance(node, VariableDefined): return ( f'{ctx.indexed(self.symbols.name[node.name], list(node.dims))} {self.format.prose(" exists")}', comparison, ) - if isinstance(node, ParameterComparisonNode): + if isinstance(node, ParameterComparison): left = ctx.indexed(self.symbols.name[node.name], list(node.dims)) return f'{left} {self._op(_PREDICATES[node.op])} {self._literal(node.value)}', comparison - if isinstance(node, DimensionComparisonNode): + if isinstance(node, DimensionComparison): if isinstance(node.value, int | float): self.noticed.numeric_coordinates.add(node.name) return ( @@ -531,38 +531,38 @@ def _where(self, node: WhereNode, ctx: _Context) -> tuple[str, int]: comparison, ) - if isinstance(node, DimensionPositionNode): + if isinstance(node, DimensionPosition): grouping = None if node.by is None else self._lookup(node.by, ctx.subscript(node.name)) place = self._position(ctx.subscript(node.name), grouping) ordinal = self._ordinal(node.name, node.position, grouping) return f'{place} {self._op(_PREDICATES[node.op])} {ordinal}', comparison - if isinstance(node, LookupComparisonNode): + if isinstance(node, LookupComparison): applied = self._lookup(node.name, ctx.subscript(node.over)) return f'{applied} {self._op(_PREDICATES[node.op])} {self._literal(node.value)}', comparison - if isinstance(node, LookupPairComparisonNode): + if isinstance(node, LookupPairComparison): index = ctx.subscript(node.over) left = self._lookup(node.name, index) right = self._lookup(node.other, index) return f'{left} {self._op(_PREDICATES[node.op])} {right}', comparison - if isinstance(node, LookupDefinedNode): + if isinstance(node, LookupDefined): applied = self._lookup(node.name, ctx.subscript(node.over)) return f'{applied} {self.format.prose(" is defined")}', comparison - if isinstance(node, NotNode): + if isinstance(node, Not): return ( f'{self._op("not")} {self._predicate(node.operand, ctx, need=_WHERE_PRECEDENCE["not"])}', _WHERE_PRECEDENCE['not'], ) - if isinstance(node, AndNode): + if isinstance(node, And): need = _WHERE_PRECEDENCE['and'] sides = [self._predicate(node.left, ctx, need=need), self._predicate(node.right, ctx, need=need)] return self.format.joined(sides, self._op('and')), need - if isinstance(node, OrNode): + if isinstance(node, Or): need = _WHERE_PRECEDENCE['or'] sides = [self._predicate(node.left, ctx, need=need), self._predicate(node.right, ctx, need=need)] return self.format.joined(sides, self._op('or')), need diff --git a/src/math_spec/validation.py b/src/math_spec/validation.py index f8316600..54e8da16 100644 --- a/src/math_spec/validation.py +++ b/src/math_spec/validation.py @@ -37,7 +37,7 @@ from math_spec.expansion import expand, parse_and_expand, parse_template from math_spec.model import Spec from math_spec.operators import BUILTINS, unknown_operator_message -from math_spec.program import BooleanLiteralNode +from math_spec.program import BooleanLiteral from math_spec.resolution import ( Namespace, Resolved, @@ -52,7 +52,7 @@ from pathlib import Path from math_spec.model import ExpressionBlock - from math_spec.program import WhereNode + from math_spec.program import Predicate def to_spec(model: str | Path | dict[str, Any] | Spec) -> Spec: @@ -185,11 +185,11 @@ def _named( found = len(errors) arms: list[CaseArm] = [] - masks: dict[str, WhereNode] = {} + masks: dict[str, Predicate] = {} for case_name, case in block.cases.items(): arm_context = case_context(name, case_name) when = resolve_where_text(case.when, ns, arm_context, errors) - if isinstance(when, BooleanLiteralNode): + if isinstance(when, BooleanLiteral): errors.append(_constant_arm(arm_context, value=when.value)) elif when is not None: masks[case_name] = when diff --git a/tests/test_dimensions.py b/tests/test_dimensions.py index a14c87fb..2c50d748 100644 --- a/tests/test_dimensions.py +++ b/tests/test_dimensions.py @@ -11,7 +11,7 @@ import pytest from math_spec.dimensions import DimensionError, _check_where_dims, dims_of -from math_spec.program import LookupPairComparisonNode, Mask +from math_spec.program import LookupPairComparison, Mask from math_spec.resolution import Namespace, expression_of, where_of from math_spec.validation import to_spec from tests.fixtures import override, schema_of @@ -392,6 +392,6 @@ def test_names_read_takes_both_sides_of_a_lookup_pair(): BASE has one lookup per dimension, so the pair is built directly rather than resolved from a predicate string. """ - where = LookupPairComparisonNode('from_bus', 'to_bus', 'line', '!=') + where = LookupPairComparison('from_bus', 'to_bus', 'line', '!=') assert Mask(where).names_read == {'from_bus', 'to_bus'}, 'a lookup pair names both maps it compares' diff --git a/tests/test_exclusivity.py b/tests/test_exclusivity.py index b91cad3f..c2f86325 100644 --- a/tests/test_exclusivity.py +++ b/tests/test_exclusivity.py @@ -19,13 +19,13 @@ from math_spec._where_parser import parse_where from math_spec.exclusivity import CELL_BUDGET, Special, Subject, _evaluate, _Grid, overlapping -from math_spec.program import AndNode, Mask, NotNode, OrNode +from math_spec.program import And, Mask, Not, Or from math_spec.resolution import Namespace, resolve_where from math_spec.validation import to_spec if TYPE_CHECKING: from math_spec.model import Spec - from math_spec.program import WhereNode + from math_spec.program import Predicate #: A storage model carrying one atom of every kind a `when` can be built from. #: Every axis takes its coordinates from data, so nothing here sizes one. @@ -60,7 +60,7 @@ def refusals(schema: Spec, cases: dict[str, str]) -> list[str]: return list(overlapping({name: _mask(when, namespace, name) for name, when in cases.items()}, namespace.dtypes)) -def _mask(text: str, namespace: Namespace, name: str) -> WhereNode: +def _mask(text: str, namespace: Namespace, name: str) -> Predicate: """Resolved but not folded, which is the shape a case's `when` reaches the prover in.""" errors: list[str] = [] mask = resolve_where(parse_where(text), namespace, f"case '{name}'", errors) @@ -252,11 +252,11 @@ class TestSoundness: def _random_mask(self, rng: random.Random, atoms: list[Any], depth: int = 0) -> Any: if depth >= 2 or rng.random() < 0.45: atom = rng.choice(atoms) - return NotNode(atom) if rng.random() < 0.25 else atom + return Not(atom) if rng.random() < 0.25 else atom left = self._random_mask(rng, atoms, depth + 1) right = self._random_mask(rng, atoms, depth + 1) - node = AndNode(left, right) if rng.random() < 0.5 else OrNode(left, right) - return NotNode(node) if rng.random() < 0.15 else node + node = And(left, right) if rng.random() < 0.5 else Or(left, right) + return Not(node) if rng.random() < 0.15 else node @pytest.mark.parametrize('seed', [1, 7]) def test_a_pair_proved_apart_stays_apart_on_a_finer_grid(self, schema: Spec, seed: int): diff --git a/tests/test_lowering.py b/tests/test_lowering.py index f743319f..80077a29 100644 --- a/tests/test_lowering.py +++ b/tests/test_lowering.py @@ -24,34 +24,34 @@ from math_spec.program import ( QUADRATIC_POSITIONS, Add, - AndNode, - At, - BooleanLiteralNode, + And, + BooleanLiteral, Cases, Constant, - DimensionComparisonNode, + DimensionComparison, DimensionDeclaration, Divide, Dual, - ExpressionNode, + Expression, Footprint, GroupSum, LookupDeclaration, Mask, Multiply, Negate, - NotNode, - OrNode, + Not, + Or, Parameter, - ParameterComparisonNode, - ParameterDefinedNode, + ParameterComparison, + ParameterDefined, Power, Program, + Pullback, Region, Sum, Translate, Variable, - Window, + WindowSum, children, divisor_parameters, fan_in, @@ -69,7 +69,7 @@ DISPATCH_YAML = EXAMPLES / 'dispatch.yaml' #: The mask `examples/dispatch.yaml` puts on `p`, as the plan carries it. -P_MAX_POSITIVE = ParameterComparisonNode('p_max', '>', 0.0, ('generator',)) +P_MAX_POSITIVE = ParameterComparison('p_max', '>', 0.0, ('generator',)) #: One dimension, one parameter, one bounded variable and a scalar constraint: #: the smallest model that loads, for a claim about the plan's record rather @@ -141,9 +141,9 @@ def test_lower_program_structure(dispatch_program): assert c.rhs == Parameter('load') assert dispatch_program.objective.sense == 'minimize', "the program carries the language's spelling, untranslated" - assert dispatch_program.objective.expression == Sum(Variable('p') * Parameter('cost'), ('generator', 'snapshot')), ( - 'the objective carries the sum the file wrote, over the dims it named none of' - ) + assert dispatch_program.objective.expression == Sum( + Multiply(Variable('p'), Parameter('cost')), ('generator', 'snapshot') + ), 'the objective carries the sum the file wrote, over the dims it named none of' @pytest.mark.parametrize('sense', [pytest.param('minimize', id='minimize'), pytest.param('maximize', id='maximize')]) @@ -173,33 +173,33 @@ def test_a_literal_amount_resolves_to_one_signed_number(dispatch_schema): [ pytest.param(None, None, id='no-where-at-all'), pytest.param('True', None, id='True-is-no-mask'), - pytest.param('p_max', ParameterDefinedNode('p_max', ('generator',)), id='a-bare-parameter-name'), + pytest.param('p_max', ParameterDefined('p_max', ('generator',)), id='a-bare-parameter-name'), pytest.param( 'snapshot > 5', - DimensionComparisonNode('snapshot', '>', 5), + DimensionComparison('snapshot', '>', 5), id='a-dimension-coordinate-compares-like-a-parameter', ), pytest.param( 'p_max > 0 AND NOT load == 0', - AndNode(P_MAX_POSITIVE, NotNode(ParameterComparisonNode('load', '==', 0.0, ('snapshot',)))), + And(P_MAX_POSITIVE, Not(ParameterComparison('load', '==', 0.0, ('snapshot',)))), id='a-compound-where-keeps-its-connectives', ), - pytest.param('False', BooleanLiteralNode(False), id='the-empty-declaration-keeps-its-own-spelling'), + pytest.param('False', BooleanLiteral(False), id='the-empty-declaration-keeps-its-own-spelling'), pytest.param('p_max > 0 AND True', P_MAX_POSITIVE, id='and-true-is-the-other-side'), pytest.param('p_max > 0 OR False', P_MAX_POSITIVE, id='or-false-is-the-other-side'), pytest.param('p_max > 0 OR True', None, id='or-true-is-no-mask-at-all'), - pytest.param('p_max > 0 AND False', BooleanLiteralNode(False), id='and-false-is-the-empty-declaration'), - pytest.param('NOT True', BooleanLiteralNode(False), id='not-true-is-false'), + pytest.param('p_max > 0 AND False', BooleanLiteral(False), id='and-false-is-the-empty-declaration'), + pytest.param('NOT True', BooleanLiteral(False), id='not-true-is-false'), pytest.param('NOT False', None, id='not-false-is-no-mask'), pytest.param('NOT (p_max > 0 AND False)', None, id='a-branch-folded-away-folds-the-one-above-it'), pytest.param( 'NOT (NOT p_max)', - ParameterDefinedNode('p_max', ('generator',)), + ParameterDefined('p_max', ('generator',)), id='a-double-negation-cancels-on-the-load-path', ), pytest.param( '(p_max > 0 OR True) AND load', - ParameterDefinedNode('load', ('snapshot',)), + ParameterDefined('load', ('snapshot',)), id='an-absorbed-side-takes-its-own-branch-with-it', ), ], @@ -207,7 +207,7 @@ def test_a_literal_amount_resolves_to_one_signed_number(dispatch_schema): def test_a_where_is_one_resolved_predicate_with_every_literal_folded(dispatch_schema, where, expected): """One mask had two lowerings: `True` was dropped at the root and kept under a connective. - A `BooleanLiteralNode` is a node a consumer meets at the root or nowhere. + A `BooleanLiteral` is a node a consumer meets at the root or nowhere. """ mask = where_of(where, Namespace.of(dispatch_schema), 't') assert (mask.root if mask is not None else None) == expected, ( @@ -221,7 +221,7 @@ def test_a_folded_mask_reaches_the_declaration_the_shorter_spelling_would_have() expand_piecewise(schema_of(DISPATCH_MODEL, **{'variables.p.where': 'p_max > 0 AND True'})) ) plain = lower_program(expand_piecewise(schema_of(DISPATCH_MODEL, **{'variables.p.where': 'p_max > 0'}))) - assert written_out.variable('p') == plain.variable('p'), 'the same mask, so the same declaration' + assert written_out.variables['p'] == plain.variables['p'], 'the same mask, so the same declaration' def test_an_unknown_where_name_is_an_error_at_lowering_too(dispatch_schema): @@ -253,9 +253,8 @@ def test_a_lowered_mask_cannot_be_rewritten_in_place(dispatch_program): def test_a_lowered_where_is_a_mask_that_answers_from_its_root(dispatch_program): """The `where` a lowering carries is a `Mask`, and its questions are its root's. - A consumer asks the mask — `where.names_read`, `where.conjuncts` — the way it - asks a dimension `dimension.targets`, rather than reaching for a free function - with the raw node. + A consumer asks the mask — `where.names_read`, `where.conjuncts` — rather + than reaching for a free function with the raw node. """ (v,) = dispatch_program.variables.values() @@ -297,17 +296,17 @@ def test_a_lowered_mask_answers_its_dims_conjuncts_and_atoms(variable, where, di assert len(mask.atoms) == atoms, 'the leaves of every arm, connectives removed' -FLAG = ParameterDefinedNode('flag', ('generator',)) +FLAG = ParameterDefined('flag', ('generator',)) @pytest.mark.parametrize( ('where', 'under'), [ - pytest.param(NotNode(P_MAX_POSITIVE), (P_MAX_POSITIVE,), id='a-not-carries-its-operand'), - pytest.param(AndNode(P_MAX_POSITIVE, FLAG), (P_MAX_POSITIVE, FLAG), id='an-and-carries-both-sides'), - pytest.param(OrNode(P_MAX_POSITIVE, FLAG), (P_MAX_POSITIVE, FLAG), id='an-or-carries-both-sides'), + pytest.param(Not(P_MAX_POSITIVE), (P_MAX_POSITIVE,), id='a-not-carries-its-operand'), + pytest.param(And(P_MAX_POSITIVE, FLAG), (P_MAX_POSITIVE, FLAG), id='an-and-carries-both-sides'), + pytest.param(Or(P_MAX_POSITIVE, FLAG), (P_MAX_POSITIVE, FLAG), id='an-or-carries-both-sides'), pytest.param(P_MAX_POSITIVE, (), id='a-leaf-carries-nothing'), - pytest.param(BooleanLiteralNode(False), (), id='a-literal-carries-nothing'), + pytest.param(BooleanLiteral(False), (), id='a-literal-carries-nothing'), ], ) def test_where_children_is_the_one_walk_under_a_predicate(where, under): @@ -325,32 +324,32 @@ def test_where_children_is_the_one_walk_under_a_predicate(where, under): def test_a_synthetic_predicate_answers_its_own_dims(): """A tree built from resolved pieces answers like a declaration's own mask. - A consumer builds region complements and conjunctions — `NotNode(root)`, - `AndNode(a, b)` — with no declaration behind them. Because the leaves carry + A consumer builds region complements and conjunctions — `Not(root)`, + `And(a, b)` — with no declaration behind them. Because the leaves carry their dims, wrapping any such tree in `Mask` answers without a name-to-dims mapping, which is what let the mapping die everywhere. """ - b = ParameterDefinedNode('load', ('snapshot',)) + b = ParameterDefined('load', ('snapshot',)) - assert Mask(NotNode(P_MAX_POSITIVE)).dims == {'generator'}, 'negation keeps the dims it negates' + assert Mask(Not(P_MAX_POSITIVE)).dims == {'generator'}, 'negation keeps the dims it negates' assert (Mask(P_MAX_POSITIVE) & Mask(b)).dims == {'generator', 'snapshot'}, 'conjunction unions both sides' - assert (Mask(P_MAX_POSITIVE) & Mask(b)).root == AndNode(P_MAX_POSITIVE, b), ( + assert (Mask(P_MAX_POSITIVE) & Mask(b)).root == And(P_MAX_POSITIVE, b), ( 'the conjunction joins the roots under one AND' ) def test_mask_construction_folds_so_a_literal_stands_at_the_root_or_nowhere(): """The fold lives in the constructor, so the invariant holds however a mask is built.""" - x = ParameterDefinedNode('committable', ('g',)) - empty, every = Mask(BooleanLiteralNode(False)), Mask(BooleanLiteralNode(True)) + x = ParameterDefined('committable', ('g',)) + empty, every = Mask(BooleanLiteral(False)), Mask(BooleanLiteral(True)) - assert Mask(OrNode(BooleanLiteralNode(True), x)) == every, 'a True side absorbs the OR at the door' - assert Mask(AndNode(BooleanLiteralNode(False), x)) == empty, 'a False side dominates the AND at the door' - assert Mask(NotNode(BooleanLiteralNode(True))) == empty, 'NOT over a literal flips at the door' - assert Mask(NotNode(NotNode(x))) == Mask(x), 'a double negation cancels at the door' + assert Mask(Or(BooleanLiteral(True), x)) == every, 'a True side absorbs the OR at the door' + assert Mask(And(BooleanLiteral(False), x)) == empty, 'a False side dominates the AND at the door' + assert Mask(Not(BooleanLiteral(True))) == empty, 'NOT over a literal flips at the door' + assert Mask(Not(Not(x))) == Mask(x), 'a double negation cancels at the door' - assert ~Mask(x) == Mask(NotNode(x)), 'a plain predicate negated gains one NOT' - assert ~Mask(NotNode(x)) == Mask(x), '`not (not x)` cancels rather than stacking, so no consumer evaluates it twice' + assert ~Mask(x) == Mask(Not(x)), 'a plain predicate negated gains one NOT' + assert ~Mask(Not(x)) == Mask(x), '`not (not x)` cancels rather than stacking, so no consumer evaluates it twice' assert ~empty == every, 'the empty mask negated admits every row, with no NOT stacked' assert ~every == empty, 'and back again' assert empty & Mask(x) == empty, 'a False root dominates the conjunction' @@ -362,9 +361,9 @@ def test_mask_construction_folds_so_a_literal_stands_at_the_root_or_nowhere(): def test_a_held_leaf_walk_is_taken_after_the_fold_absorbed_a_branch(): """`atoms` is held from construction, and construction folds first — so the fold's losses are not in it.""" - absorbed = Mask(AndNode(BooleanLiteralNode(False), ParameterDefinedNode('committable', ('g',)))) + absorbed = Mask(And(BooleanLiteral(False), ParameterDefined('committable', ('g',)))) - assert absorbed.root == BooleanLiteralNode(False) + assert absorbed.root == BooleanLiteral(False) assert absorbed.atoms == (), 'the absorbed leaf is not among them' assert absorbed.names_read == frozenset(), 'nor named' assert absorbed.dims == frozenset(), 'nor read at any dim' @@ -389,7 +388,7 @@ def test_a_constraint_where_is_a_mask_like_a_variable_s(): lowered = to_program(override(DISPATCH_MODEL, **{'constraints.balance.where': 'load > 0'})) (c,) = lowered.constraints.values() - assert c.where == Mask(ParameterComparisonNode('load', '>', 0.0, ('snapshot',))) + assert c.where == Mask(ParameterComparison('load', '>', 0.0, ('snapshot',))) def test_a_power_lowers_to_a_node_of_its_own(dispatch_schema): @@ -419,7 +418,7 @@ def test_a_power_lowers_to_a_node_of_its_own(dispatch_schema): ), pytest.param( 'at(r, by=lk)', - At(Variable('r'), over='g', coordinate=('lk',), into=('h',)), + Pullback(Variable('r'), over='g', coordinate=('lk',), into=('h',)), id='a-pullback-walks-the-same-table-back', ), pytest.param( @@ -444,17 +443,17 @@ def test_a_power_lowers_to_a_node_of_its_own(dispatch_schema): ), pytest.param( 'sum_back(p, over=g, within=3)', - Window(Variable('p'), 'g', width=3, wrap=False), + WindowSum(Variable('p'), 'g', width=3, wrap=False), id='a-window-is-one-node-rather-than-a-fold-of-translations', ), pytest.param( 'sum_back(p, over=g, within=k)', - Window(Variable('p'), 'g', width='k', wrap=False), + WindowSum(Variable('p'), 'g', width='k', wrap=False), id='a-named-width-crosses-as-the-parameter-name', ), pytest.param( 'sum_back(p, over=g, within=2, by=lk)', - Window(Variable('p'), 'g', width=2, wrap=False, partition='lk'), + WindowSum(Variable('p'), 'g', width=2, wrap=False, partition='lk'), id='a-window-stops-at-the-edges-of-the-lookup-it-names', ), ], @@ -467,15 +466,15 @@ def test_a_construct_lowers_to_its_node(shapes_schema, expression, expected): def test_a_binary_variable_lowers_to_a_binary_domain(): program = to_program(schema_of(DISPATCH_YAML, **{'variables.p.domain': 'binary', 'variables.p.bounds': {}})) - assert program.variable('p').domain == 'binary' + assert program.variables['p'].domain == 'binary' def test_a_divisor_under_a_pullback_is_still_named(): """`children` has to descend through every node, or a refusal loses its name.""" quotient = Divide(Variable('x'), Parameter('rate')) - pulled = At(quotient, over='flow', coordinate=('component',), into=('component',)) + pulled = Pullback(quotient, over='flow', coordinate=('component',), into=('component',)) - assert divisor_parameters(pulled) == frozenset({'rate'}), 'the walk descends through `At`' + assert divisor_parameters(pulled) == frozenset({'rate'}), 'the walk descends through `Pullback`' assert divisor_parameters(Sum(pulled, ('flow',))) == frozenset({'rate'}), 'and through a `Sum` over it' @@ -493,12 +492,12 @@ def test_a_quotient_is_found_whole_so_its_two_halves_stay_paired(): left = Divide(Variable('x'), Parameter('rate')) right = Divide(Variable('y'), Parameter('loss')) - found = quotients(Sum(left + right, ('flow',))) + found = quotients(Sum(Add(left, right), ('flow',))) assert [(variables_of(q.numerator), q.divisor) for q in found] == [ (frozenset({'x'}), Parameter('rate')), (frozenset({'y'}), Parameter('loss')), ], 'each quotient keeps its own numerator, in the order the expression writes them' - assert divisor_parameters(Sum(left + right, ('flow',))) == frozenset({'rate', 'loss'}), ( + assert divisor_parameters(Sum(Add(left, right), ('flow',))) == frozenset({'rate', 'loss'}), ( 'the flat answer is still the union of the same walk' ) @@ -514,10 +513,10 @@ def test_a_quotient_is_found_whole_so_its_two_halves_stay_paired(): Divide(Variable('p'), Parameter('c')): 'one-to-one', Sum(Variable('p'), ('g',)): 'many-to-one', GroupSum(Variable('p'), over='g', coordinate=('at_bus',), into=('bus',)): 'many-to-one', - At(Variable('p'), over='g', coordinate=('at_bus',), into=('bus',)): 'one-to-one', + Pullback(Variable('p'), over='g', coordinate=('at_bus',), into=('bus',)): 'one-to-one', Translate(Variable('p'), 't', offset=1, wrap=False, fill=0.0): 'one-to-one', - Window(Variable('p'), 't', width=2, wrap=False): 'one-to-many', - Cases((Region(Mask(ParameterDefinedNode('c', ('g',))), Variable('p')),)): 'one-to-one', + WindowSum(Variable('p'), 't', width=2, wrap=False): 'one-to-many', + Cases((Region(Mask(ParameterDefined('c', ('g',))), Variable('p')),)): 'one-to-one', Dual('balance'): 'one-to-one', } @@ -525,8 +524,8 @@ def test_a_quotient_is_found_whole_so_its_two_halves_stay_paired(): def test_every_expression_node_is_classified_by_fan_in(): """`fan_in` was a ClassVar on five nodes, so `Add(...).fan_in` was an AttributeError.""" covered = {type(node) for node in FAN_IN} - assert covered == set(get_args(ExpressionNode)), ( - 'every node in the ExpressionNode union is classified, and nothing retired lingers' + assert covered == set(get_args(Expression)), ( + 'every node in the Expression union is classified, and nothing retired lingers' ) @@ -535,37 +534,21 @@ def test_a_node_answers_its_fan_in(node, expected): assert fan_in(node) == expected -def test_a_lookup_names_the_dimension_its_values_label(): - """Five sites asked this and each walked for it; the plan answers it once.""" - program = Program( - parameters={}, - variables={}, - constraints={}, - objective=None, - dimensions={ - 'snapshot': DimensionDeclaration((LookupDeclaration('season_of', 'season'),)), - 'generator': DimensionDeclaration((LookupDeclaration('at_bus', 'bus'),)), - }, - ) - - assert program.dimension('snapshot').targets == {'season_of': 'season'}, ( - 'one dimension names its own maps and no other dimension' +def test_a_lookup_is_declared_as_the_file_declares_it(): + """One group keyed by name, each entry naming the dimension it leaves and the one it lands in.""" + program = to_program( + override( + TINY, + dimensions={'g': {}, 'bus': {}, 'season': {}}, + lookups={'at_bus': {'over': 'g', 'into': 'bus'}, 'season_of': {'over': 'g', 'into': 'season'}}, + ) ) - assert program.dimension('generator').targets == {'at_bus': 'bus'}, 'and the same for the second' - assert [(d, lk.name) for d, lk in program.lookups] == [ - ('snapshot', 'season_of'), - ('generator', 'at_bus'), - ], 'every map with the dimension it is over, in declaration order' - -def test_an_unknown_dimension_is_a_near_miss_rather_than_an_empty_declaration(): - """A typo used to return an empty declaration, silently dropping every join.""" - program = to_program(override(TINY, **{'dimensions.snapshot': {}})) - - assert program.dimension('snapshot').dtype == 'str', 'a declared dimension still comes back' - with pytest.raises(KeyError, match='snapshto') as excinfo: - program.dimension('snapshto') - assert 'snapshot' in str(excinfo.value), 'the message names the near miss, which is the whole point of raising' + assert program.lookups == { + 'at_bus': LookupDeclaration(over='g', into='bus'), + 'season_of': LookupDeclaration(over='g', into='season'), + }, 'every lookup under its own name, in declaration order, and nothing nested under a dimension' + assert program.dimensions['g'] == DimensionDeclaration(dtype='str'), 'a dimension carries its dtype and no map' def test_a_program_is_built_by_keyword_so_a_field_added_later_cannot_reorder_an_old_call(): @@ -581,8 +564,8 @@ def test_a_program_seals_its_declaration_groups(dispatch_program, group): getattr(dispatch_program, group)['sneak'] = None # pyrefly: ignore[unsupported-operation] the point of the test -def test_expressions_are_the_ones_a_row_is_built_from(): - """`expressions` named the *declared* ones, which build no row at all.""" +def test_roots_are_the_trees_a_row_is_built_from(): + """`expressions` is the file's own section, which builds no row at all; the row-building trees are `roots`.""" program = to_program( override( TINY, @@ -591,16 +574,16 @@ def test_expressions_are_the_ones_a_row_is_built_from(): ) ) - assert list(program.named_expressions) == ['spend'], 'the declared ones keep their own name' - assert program.expressions == ( + assert list(program.expressions) == ['spend'], 'the declared ones keep their own name' + assert program.roots == ( program.objective.expression, program.constraints['c'].lhs, program.constraints['c'].rhs, ), 'the objective first, then both sides of each constraint, in declaration order' - assert program.named_expressions['spend'].expression not in program.expressions, ( + assert program.expressions['spend'].expression not in program.roots, ( 'a named expression builds no row, so it is not one of the expressions a row is built from' ) - assert len(program.expressions) == 3, 'and nothing else is counted' + assert len(program.roots) == 3, 'and nothing else is counted' def _footprint_of(constraint: str, objective: str) -> Footprint: @@ -639,7 +622,7 @@ def test_a_construct_the_file_does_not_use_is_an_empty_set_rather_than_none(): assert footprint.sos_types == frozenset(), 'a file declaring no sos' assert footprint.quadratic == frozenset(), 'a file with no quadratic anywhere' assert footprint.domains == {'continuous'}, 'never empty — a program has variables' - assert {type(f) for f in (footprint.sos_types, footprint.quadratic, footprint.shapes)} == {frozenset}, ( + assert {type(f) for f in (footprint.sos_types, footprint.quadratic, footprint.nodes)} == {frozenset}, ( 'every field is a set, so one rule reads all of them' ) assert footprint.quadratic <= QUADRATIC_POSITIONS, 'and the vocabulary a consumer pins its table against' @@ -657,8 +640,8 @@ def test_a_named_expression_is_not_in_the_footprint(): """It builds no row, so counting it would answer wrongly about what is solved.""" program = to_program(override(TINY, expressions={'spend': 'sum(p * cost, over=g)'})) - assert Parameter not in program.footprint.shapes, "the named expression's parameter reaches no row" - assert Parameter in {type(n) for n in walk(program.named_expressions['spend'].expression)}, ( + assert Parameter not in program.footprint.nodes, "the named expression's parameter reaches no row" + assert Parameter in {type(n) for n in walk(program.expressions['spend'].expression)}, ( 'though it is in the expression' ) @@ -671,8 +654,8 @@ def test_a_dimension_carries_the_dtype_its_labels_are_checked_against(): """ program = to_program(override(TINY, **{'dimensions.t': {'dtype': 'int'}})) - assert program.dimension('t').dtype == 'int', 'a declared dtype reaches the plan' - assert program.dimension('g').dtype == 'str', "and the schema's default does too, rather than nothing" + assert program.dimensions['t'].dtype == 'int', 'a declared dtype reaches the plan' + assert program.dimensions['g'].dtype == 'str', "and the schema's default does too, rather than nothing" CASED = { @@ -719,10 +702,8 @@ def test_the_fallback_region_carries_the_mask_the_file_left_unwritten(): """ remainder = _cases_in(to_program(CASED)).regions[-1] - assert isinstance(remainder.when.root, AndNode), ( - 'two stated cases, so the remainder is a conjunction of two negations' - ) - assert remainder.when.root.left == ParameterDefinedNode('committable', ('g',)), ( + assert isinstance(remainder.when.root, And), 'two stated cases, so the remainder is a conjunction of two negations' + assert remainder.when.root.left == ParameterDefined('committable', ('g',)), ( 'the negation of `not committable` is the term itself, not a second `not` around it' ) @@ -759,10 +740,10 @@ def test_the_lowered_regions_are_still_proved_apart(): def test_a_cased_expression_is_readable_by_the_name_the_file_wrote(): - """`Program.expressions` carries it, so a consumer reads it back whole.""" + """`Program.expressions` carries it under its name, so a consumer reads it back whole.""" program = to_program(CASED) - assert isinstance(program.named_expressions['previous'].expression, Cases), ( + assert isinstance(program.expressions['previous'].expression, Cases), ( 'a cased expression reaches the program as the node, not as its fallback arm alone' ) @@ -794,7 +775,7 @@ def test_a_cased_expression_is_readable_by_the_name_the_file_wrote(): def test_an_entry_is_in_the_math_where_the_objective_or_a_constraint_inlines_it(patch, in_math): """`in_math` is usage, not shape: one affine body is in the math when a row inlines it, however indirectly, and a reported quantity when none does.""" program = to_program(override(TINY, expressions={'spend': 'sum(p * cost, over=g)'}, **patch)) - assert program.named_expressions['spend'].in_math is in_math + assert program.expressions['spend'].in_math is in_math def test_an_entry_reached_only_through_another_is_in_the_math_with_it(): @@ -806,7 +787,7 @@ def test_an_entry_reached_only_through_another_is_in_the_math_with_it(): **{'constraints.c.expression': 'twice >= 1'}, ) ) - reads = {name: program.named_expressions[name].in_math for name in ('twice', 'spend')} + reads = {name: program.expressions[name].in_math for name in ('twice', 'spend')} assert reads == {'twice': True, 'spend': True}, ( 'the entry the row names and the one it reaches through are both in the math' ) @@ -822,7 +803,7 @@ def test_a_macro_formal_named_like_an_entry_keeps_the_entry_out_of_the_math(): **{'constraints.c.expression': 'scaled(sum(p, over=g)) >= 1'}, ) ) - assert program.named_expressions['spend'].in_math is False, ( + assert program.expressions['spend'].in_math is False, ( 'the formal shadows the entry, so the constraint inlines the argument and the math never reads spend' ) @@ -830,7 +811,7 @@ def test_a_macro_formal_named_like_an_entry_keeps_the_entry_out_of_the_math(): def test_an_entry_that_reads_a_dual_is_a_reported_quantity(): """A dual is read after the solve, so an entry calling one is never in the math: it lowers to a Dual leaf and stays reported.""" program = to_program(override(TINY, expressions={'shadow_price': 'dual(c)'})) - declaration = program.named_expressions['shadow_price'] + declaration = program.expressions['shadow_price'] assert declaration.in_math is False, 'the entry reading a dual is reported, never in the math' assert isinstance(declaration.expression, Dual), 'and it lowers to a Dual leaf' diff --git a/tests/test_parser.py b/tests/test_parser.py index 3748a75f..74160a70 100644 --- a/tests/test_parser.py +++ b/tests/test_parser.py @@ -32,7 +32,7 @@ parse_where, ) from math_spec.errors import SchemaError -from math_spec.program import AndNode, BooleanLiteralNode, NotNode, OrNode, _conjuncts +from math_spec.program import And, BooleanLiteral, Not, Or, _conjuncts def test_the_grammar_builds_the_program_s_own_node_classes(): @@ -40,13 +40,13 @@ def test_the_grammar_builds_the_program_s_own_node_classes(): The parser constructs the resolved vocabulary's connectives directly, so a consumer's `isinstance` against the program's classes holds on any tree — - two homes for `AndNode` would make it hold on neither. + two homes for `And` would make it hold on neither. """ tree = parse_where('a AND NOT b OR True') - assert type(tree) is program_module.OrNode - assert type(tree.left) is program_module.AndNode - assert type(tree.left.right) is program_module.NotNode - assert type(tree.right) is program_module.BooleanLiteralNode + assert type(tree) is program_module.Or + assert type(tree.left) is program_module.And + assert type(tree.left.right) is program_module.Not + assert type(tree.right) is program_module.BooleanLiteral @pytest.mark.parametrize( @@ -240,12 +240,12 @@ def test_a_name_may_begin_with_inf(name): @pytest.mark.parametrize( ('text', 'node_type', 'attrs'), [ - pytest.param('True', BooleanLiteralNode, {'value': True}, id='a-literal'), + pytest.param('True', BooleanLiteral, {'value': True}, id='a-literal'), pytest.param('p_max', UnresolvedNameNode, {'name': 'p_max'}, id='a-bare-name'), pytest.param('p_max > 0', UnresolvedComparisonNode, {'op': '>', 'value': 0}, id='a-comparison'), - pytest.param('a AND b', AndNode, {}, id='and'), - pytest.param('a OR b', OrNode, {}, id='or'), - pytest.param('NOT a', NotNode, {}, id='not'), + pytest.param('a AND b', And, {}, id='and'), + pytest.param('a OR b', Or, {}, id='or'), + pytest.param('NOT a', Not, {}, id='not'), ], ) def test_a_where_string_parses_to_its_node(text, node_type, attrs): @@ -258,8 +258,8 @@ def test_a_where_string_parses_to_its_node(text, node_type, attrs): def test_and_binds_tighter_than_or(): - assert parse_where('a OR b AND c') == OrNode( - UnresolvedNameNode('a'), AndNode(UnresolvedNameNode('b'), UnresolvedNameNode('c')) + assert parse_where('a OR b AND c') == Or( + UnresolvedNameNode('a'), And(UnresolvedNameNode('b'), UnresolvedNameNode('c')) ) @@ -273,7 +273,7 @@ def test_and_binds_tighter_than_or(): ids=['single', 'pair', 'chain'], ) def test_conjuncts_flattens_the_and_spine(text, expected): - """A chain the grammar left-folds into nested `AndNode`s comes back flat (#312). + """A chain the grammar left-folds into nested `And`s comes back flat (#312). `_conjuncts` is the one home of the flatten rule; `Mask.conjuncts` is the door a consumer asks it through.""" @@ -336,7 +336,7 @@ def test_position_converts_a_dimension_to_where_a_row_sits(text, op, position, b def test_a_position_is_not_confused_with_a_name(): """`position` leads the alternation, so it is not read as a bare name.""" - assert isinstance(parse_where('position(t) == 0 AND p_max > 0'), AndNode) + assert isinstance(parse_where('position(t) == 0 AND p_max > 0'), And) @pytest.mark.parametrize( diff --git a/tests/test_piecewise.py b/tests/test_piecewise.py index 57e3cfd6..50305cca 100644 --- a/tests/test_piecewise.py +++ b/tests/test_piecewise.py @@ -321,7 +321,7 @@ def test_an_entry_a_link_reads_is_in_the_math(): NONCONVEX_YAML, **{'expressions': {'twice': 'p * 2'}, 'piecewise.cost_curve.links': [['twice', 'bp_x'], ['op_cost', 'bp_y']]}, ) - assert to_program(schema).named_expressions['twice'].in_math is True + assert to_program(schema).expressions['twice'].in_math is True def test_a_link_reading_a_dual_entry_is_refused(): diff --git a/tests/test_program_nodes.py b/tests/test_program_nodes.py index 030b813a..34f7f7bc 100644 --- a/tests/test_program_nodes.py +++ b/tests/test_program_nodes.py @@ -7,7 +7,7 @@ The lowering-side sibling of `test_the_golden_model_carries_every_node_kind_the_walk_renders`, on a fixture of its own because rendering accepts what lowering refuses. -Without this, a node can join `ExpressionNode` with nothing producing it and +Without this, a node can join `Expression` with nothing producing it and the suite stays green — `assert_never` fires only where some test happens to lower a file that uses the construct. That is how `cases:` reached a release candidate unlowerable. @@ -21,12 +21,12 @@ import pytest import math_spec as ms -from math_spec.program import ExpressionNode, Program, walk +from math_spec.program import Expression, Program, walk FIXTURE = Path(__file__).resolve().parent / 'fixtures' / 'every_program_node.yaml' -def _expressions(program: Program) -> list[ExpressionNode]: +def _expressions(program: Program) -> list[Expression]: """Every tree a program hangs on to, wherever it hangs it. Bounds and named expressions among them: a node reachable only from an @@ -35,7 +35,7 @@ def _expressions(program: Program) -> list[ExpressionNode]: """ trees = [side for c in program.constraints.values() for side in (c.lhs, c.rhs)] trees += [bound for v in program.variables.values() for bound in (v.lower, v.upper)] - trees += [e.expression for e in program.named_expressions.values()] + trees += [e.expression for e in program.expressions.values()] if program.objective is not None: trees.append(program.objective.expression) return trees @@ -43,10 +43,10 @@ def _expressions(program: Program) -> list[ExpressionNode]: @pytest.fixture(scope='module') def kinds() -> tuple[set[str], set[str]]: - """The node classes the fixture lowers to, and the ones `ExpressionNode` declares.""" + """The node classes the fixture lowers to, and the ones `Expression` declares.""" program = ms.to_program(FIXTURE) reached = {type(node).__name__ for node in walk(*_expressions(program))} - declared = {node.__name__ for node in get_args(ExpressionNode)} + declared = {node.__name__ for node in get_args(Expression)} return reached, declared diff --git a/tests/test_validation.py b/tests/test_validation.py index 2fcc591a..3ea80c1d 100644 --- a/tests/test_validation.py +++ b/tests/test_validation.py @@ -15,7 +15,7 @@ from math_spec._yaml import parse_yaml from math_spec.errors import DimensionError, LanguageError, SchemaError from math_spec.lowering import to_program -from math_spec.program import DimensionPositionNode +from math_spec.program import DimensionPosition from math_spec.resolution import Namespace, where_of from math_spec.typesetting import to_markdown from math_spec.validation import to_spec @@ -176,7 +176,7 @@ def test_an_unreferenced_nonlinear_entry_loads_and_is_reported(self): nothing consumes. """ model = override(SMALL_MODEL, expressions={'lcoe': 'c / sum(p)'}) - assert to_program(model).named_expressions['lcoe'].in_math is False, ( + assert to_program(model).expressions['lcoe'].in_math is False, ( 'the unread nonlinear body loads rather than being refused, and nothing in the math reads it' ) assert 'lcoe' in to_markdown(model), 'and the page prints it, under its own name' @@ -457,7 +457,7 @@ def test_it_resolves(self, mask: str, position: int, by: str | None): resolved = where_of(mask, Namespace.of(POSITION_SCHEMA), 'the mask') assert resolved is not None node = resolved.root - assert isinstance(node, DimensionPositionNode) + assert isinstance(node, DimensionPosition) assert node.name == 'snapshot' assert node.position == position assert node.by == by diff --git a/tests/typesetting/test_golden.py b/tests/typesetting/test_golden.py index e93b0712..f0a1c983 100644 --- a/tests/typesetting/test_golden.py +++ b/tests/typesetting/test_golden.py @@ -18,7 +18,7 @@ from math_spec._expression_parser import ArithmeticNode, ComparisonNode, DualNode, FunctionCallNode from math_spec.operators import BUILTIN_NAMES from math_spec.piecewise import expand_piecewise -from math_spec.program import WhereNode +from math_spec.program import Predicate from math_spec.typesetting import FORMATS, to_latex, typeset, walk from math_spec.typesetting.format import OPERATOR_NAMES from math_spec.validation import to_spec @@ -161,7 +161,7 @@ def test_the_golden_model_carries_every_node_kind_the_walk_renders(): `coverage` installed, and its failure names the construct rather than a line. """ kinds = {type(node).__name__ for tree in _rendered_trees() for node in _nodes(tree)} - CARRIERS - declared = {node.__name__ for node in (*get_args(WhereNode), *get_args(ArithmeticNode), ComparisonNode)} + declared = {node.__name__ for node in (*get_args(Predicate), *get_args(ArithmeticNode), ComparisonNode)} assert kinds == declared - UNRESOLVED, ( f'tests/typesetting/golden/model.yaml reaches {sorted(kinds - declared)} and misses ' f'{sorted(declared - UNRESOLVED - kinds)}. Every node the walk renders needs a case here, ' From f4883d2e0b6158a07bcb860b341e4d72ef066d86 Mon Sep 17 00:00:00 2001 From: Claude Date: Tue, 15 Sep 2026 20:19:28 +0000 Subject: [PATCH 2/2] refactor(program): a translation and a window call their axis along, as the file does since #477 Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_019fhGZgaBspo7mh9Hjd3KtT --- src/math_spec/program.py | 8 ++++---- src/math_spec/separability.py | 2 +- 2 files changed, 5 insertions(+), 5 deletions(-) diff --git a/src/math_spec/program.py b/src/math_spec/program.py index 43f09f25..262cb065 100644 --- a/src/math_spec/program.py +++ b/src/math_spec/program.py @@ -310,9 +310,9 @@ class Translate: and contribute it. Always ``None`` under ``wrap``. ``offset`` is an integer, or the name of an integer parameter that does - not depend on ``over`` and carries its sign in the values. + not depend on ``along`` and carries its sign in the values. - ``partition`` is a relation walked along ``over`` — its consumed + ``partition`` is a relation walked along ``along`` — 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 @@ -321,7 +321,7 @@ class Translate: """ operand: Expression - over: str + along: str offset: int | str wrap: bool fill: float | None = None @@ -350,7 +350,7 @@ class WindowSum: """ operand: Expression - over: str + along: str width: int | str wrap: bool partition: Walk | None = None diff --git a/src/math_spec/separability.py b/src/math_spec/separability.py index f8a6d16f..cf0c70a0 100644 --- a/src/math_spec/separability.py +++ b/src/math_spec/separability.py @@ -96,7 +96,7 @@ def waits_on(dimension: str, label: str, name: str, kind: Literal['offset', 'par for relation in node.coordinate: waits_on(dimension, label, relation, 'coordinate') elif isinstance(node, (Translate, WindowSum)): - dimension = node.over + dimension = node.along if node.wrap: report( 'coupled',