diff --git a/docs/contributing.md b/docs/contributing.md index 85b03af6..2a1ea510 100644 --- a/docs/contributing.md +++ b/docs/contributing.md @@ -98,14 +98,24 @@ stale anchor fails it. `pixi run docs-serve` builds the site and serves it at The same construct passes through three layers, and each names it in full. The suffix says which layer: -| 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` | - -A node names the coordinate map rather than the operator: the translation node -is `Translate`, and the operator is `shift`. Nothing is abbreviated. +| 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` | + +A node names the operation, not the verb a file writes. One verb can lower to +two nodes, so the file's spelling cannot decide the name. + +| File verb | Node | What the node names | +| ------------------ | ----------- | ------------------------------ | +| `sum(over=)` | `Sum` | dims removed from the result | +| `sum(by=)` | `GroupSum` | a sum through a relation | +| `at(by=)` | `Pullback` | a read through a relation | +| `shift(along=)` | `Translate` | a re-index along one dimension | +| `sum_back(along=)` | `WindowSum` | a sum over a trailing window | + +Nothing is abbreviated. ## Adding an operator diff --git a/docs/reference/reading.md b/docs/reference/reading.md index 7d26b4b4..ebb4bb3e 100644 --- a/docs/reference/reading.md +++ b/docs/reference/reading.md @@ -110,7 +110,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.kinds) # ['Constant', 'Multiply', 'Parameter', 'Sum', 'Variable'] ``` Every field is a set. An empty field means this model does not use the diff --git a/src/math_spec/_expression_parser.py b/src/math_spec/_expression_parser.py index 52bdf216..ba740563 100644 --- a/src/math_spec/_expression_parser.py +++ b/src/math_spec/_expression_parser.py @@ -23,7 +23,7 @@ if TYPE_CHECKING: from collections.abc import Callable, Iterator, Mapping - from math_spec.program import Direction, Partition, WhereNode + from math_spec.program import Direction, Partition, 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. @@ -222,7 +222,7 @@ class CaseArm: """ label: str - when: WhereNode | None + when: Predicate | None value: ArithmeticNode @@ -307,7 +307,7 @@ def __str__(self) -> str: #: 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 ed9b0123..eea60243 100644 --- a/src/math_spec/_where_parser.py +++ b/src/math_spec/_where_parser.py @@ -22,13 +22,13 @@ from math_spec._expression_parser import ARITHMETIC, NAME, ArithmeticNode, children, parse_text from math_spec.program import ( - AndNode, - BooleanLiteralNode, - ConnectiveWhereNode, - NotNode, - OrNode, + And, + BooleanLiteral, + Connective, + Not, + Or, + Predicate, PredicateOperator, - WhereNode, where_children, ) @@ -93,7 +93,7 @@ class UnresolvedComparisonNode: #: Every node a parsed where string is built of: the connectives and literals, #: the unresolved leaves, and the arithmetic and the two side nodes under a #: comparison. What the depth measurement walks. -_ParsedWhere = WhereNode | UnresolvedWhereNode | ArithmeticNode | ColumnNode | QuotedNode +_ParsedWhere = Predicate | UnresolvedWhereNode | ArithmeticNode | ColumnNode | QuotedNode # --------------------------------------------------------------------------- @@ -110,8 +110,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)) name = pp.Regex(NAME) # pyrefly: ignore[implicit-any-lambda] @@ -133,28 +133,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], WhereNode | UnresolvedWhereNode]: +def _folder(node_type: type[And] | type[Or]) -> Callable[[pp.ParseResults], Predicate | UnresolvedWhereNode]: """A parse action left-folding a flat operator chain into *node_type*.""" - def fold(tokens: pp.ParseResults) -> WhereNode | UnresolvedWhereNode: - items: list[WhereNode | UnresolvedWhereNode] = list(tokens) + def fold(tokens: pp.ParseResults) -> Predicate | UnresolvedWhereNode: + items: list[Predicate | UnresolvedWhereNode] = list(tokens) result = items[0] for item in items[1:]: - result = node_type(cast('WhereNode', result), cast('WhereNode', item)) + result = node_type(cast('Predicate', result), cast('Predicate', item)) return result return fold @@ -196,13 +196,13 @@ def _nested(node: _ParsedWhere) -> tuple[_ParsedWhere, ...]: return (node.left, node.right) if isinstance(node, ArithmeticNode): return children(node) - if isinstance(node, ConnectiveWhereNode): + if isinstance(node, Connective): return where_children(node) return () @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 @@ -217,6 +217,6 @@ def parse_where(text: str) -> WhereNode | UnresolvedWhereNode: as an expression is. """ return cast( - 'WhereNode | UnresolvedWhereNode', + 'Predicate | UnresolvedWhereNode', parse_text(_WHERE_GRAMMAR, text, 'where string', _named_rewrite, _nested, _DEEP_REWRITE), ) diff --git a/src/math_spec/advice.py b/src/math_spec/advice.py index f8e7a5b8..ee3406e5 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 collections.abc import Mapping @@ -73,7 +73,7 @@ def _produced_axes(program: Program) -> set[str]: dimension: either way, the dims the direction produces. """ axes: set[str] = set() - for node in walk(*program.expressions): - if isinstance(node, GroupSum | At): + for node in walk(*program.roots): + if isinstance(node, GroupSum | Pullback): axes.update(node.direction.produced_dims) 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 6bec4fbf..72db88a2 100644 --- a/src/math_spec/dimensions.py +++ b/src/math_spec/dimensions.py @@ -41,17 +41,17 @@ from math_spec.errors import DimensionError from math_spec.operators import BUILTINS from math_spec.program import ( - DimensionComparisonNode, - DimensionPositionNode, + DimensionComparison, + DimensionPosition, Direction, Mask, - ParameterComparisonNode, - ParameterDefinedNode, + ParameterComparison, + ParameterDefined, Partition, - RelationComparisonNode, - RelationDefinedNode, - RelationPairComparisonNode, - VariableDefinedNode, + RelationComparison, + RelationDefined, + RelationPairComparison, + VariableDefined, ) if TYPE_CHECKING: @@ -539,13 +539,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 RelationComparisonNode() | RelationPairComparisonNode() | RelationDefinedNode(): + case RelationComparison() | RelationPairComparison() | RelationDefined(): noun = 'relation' case _: assert_never(atom) diff --git a/src/math_spec/exclusivity.py b/src/math_spec/exclusivity.py index 4c816a3a..985fcf43 100644 --- a/src/math_spec/exclusivity.py +++ b/src/math_spec/exclusivity.py @@ -22,27 +22,27 @@ from typing import TYPE_CHECKING, Literal, assert_never from math_spec.program import ( - AndNode, - BooleanLiteralNode, - DimensionComparisonNode, - DimensionPositionNode, + And, + BooleanLiteral, + DimensionComparison, + DimensionPosition, Mask, - NotNode, - OrNode, - ParameterComparisonNode, - ParameterDefinedNode, - RelationComparisonNode, - RelationDefinedNode, - RelationPairComparisonNode, - TypedPredicateNode, - VariableDefinedNode, + Not, + Or, + ParameterComparison, + ParameterDefined, + RelationComparison, + RelationDefined, + RelationPairComparison, + 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) @@ -190,16 +190,16 @@ def witness(self, cell: dict[Subject, Cell]) -> str: def _observe( - node: TypedPredicateNode, subject: Subject, values: set[_Literal], dtypes: Mapping[str, DeclaredDtype] + node: TypedPredicate, subject: Subject, values: set[_Literal], 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, RelationPairComparisonNode): + elif isinstance(node, RelationPairComparison): if node.op not in ('==', '!='): msg = ( f'{subject} is ordered with {node.op!r}, and two relations carry no order ' @@ -207,7 +207,7 @@ def _observe( f'ordering as a boolean parameter and test that' ) raise Undecidable(msg) - elif isinstance(node, ParameterComparisonNode | DimensionComparisonNode | RelationComparisonNode): + elif isinstance(node, ParameterComparison | DimensionComparison | RelationComparison): 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 ' @@ -218,21 +218,21 @@ def _observe( 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, partition=partition): + case DimensionPosition(name=name, partition=partition): if partition is None: return Subject('rank', name) return Subject('rank', name, partition.name, partition.group) - case RelationDefinedNode(name=name) | RelationComparisonNode(name=name): + case RelationDefined(name=name) | RelationComparison(name=name): return Subject('relation', name) - case RelationPairComparisonNode(name=name, other=other): + case RelationPairComparison(name=name, other=other): return Subject('relation_pair', name, other) case _: assert_never(node) @@ -396,42 +396,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() | RelationDefinedNode(): + case ParameterDefined() | RelationDefined(): 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 RelationPairComparisonNode(op=op): + case RelationPairComparison(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) | RelationComparisonNode(op=op, value=literal): + case ParameterComparison(op=op, value=literal) | RelationComparison(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 4f18f4bb..850ef854 100644 --- a/src/math_spec/lowering.py +++ b/src/math_spec/lowering.py @@ -140,13 +140,7 @@ def lower_program(expanded: _ExpandedSpec) -> program.Program: _Lowering(expanded, 'the objective').expr(resolved.objective), ) - dimensions = { - dname: program.DimensionDeclaration( - tuple(lk for lk in resolved.relations.values() if dname in lk.dims), - ddef.dtype, - ) - for dname, ddef in expanded.dimensions.items() - } + dimensions = {dname: program.DimensionDeclaration(ddef.dtype) for dname, ddef in expanded.dimensions.items()} sos = { sname: program.SosDeclaration( sdef.variable, @@ -167,9 +161,10 @@ def lower_program(expanded: _ExpandedSpec) -> program.Program: constraints=constraints, objective=objective, dimensions=dimensions, + relations=resolved.relations, sos=sos, piecewise={name: declaration_of(ex) for name, ex in expanded.expanded_piecewise.items()}, - named_expressions=expressions, + expressions=expressions, ) @@ -185,7 +180,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) @@ -254,7 +249,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=relation)``. Two program nodes under one surface verb: reducing a dim away and reducing it @@ -272,13 +267,13 @@ def sum(self, node: FunctionCallNode) -> program.ExpressionNode: assert isinstance(by_node, DirectionNode), 'resolution reads sum(by=) in a direction' return program.GroupSum(operand, direction=by_node.direction) - def at(self, node: FunctionCallNode) -> program.ExpressionNode: + def at(self, node: FunctionCallNode) -> program.Expression: """``at(x, by=relation)`` — the adjoint of :meth:`sum`'s ``by=`` form.""" by_node = node.kwargs['by'] assert isinstance(by_node, DirectionNode), 'resolution reads at(by=) in a direction' - return program.At(self.expr(node.args[0]), direction=by_node.direction) + return program.Pullback(self.expr(node.args[0]), direction=by_node.direction) - def sum_back(self, node: FunctionCallNode) -> program.ExpressionNode: + def sum_back(self, node: FunctionCallNode) -> program.Expression: """``sum_back(x, along=d, window=w)`` — a trailing window along one dimension. *window* is an integer literal of at least one, or a parameter naming a @@ -300,9 +295,9 @@ def sum_back(self, node: FunctionCallNode) -> program.ExpressionNode: else: assert isinstance(window_node, NumberNode), 'a window= that is neither is refused at load' width = int(window_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, along=d, offset=n)`` — the value at *t - offset* along one dim. What the vacated positions contribute is ``edge=``'s to say, and the @@ -330,7 +325,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, @@ -352,7 +347,7 @@ def _partition_of(node: FunctionCallNode) -> program.Partition | None: return by_node.partition -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 d718445e..a3d1aa94 100644 --- a/src/math_spec/program.py +++ b/src/math_spec/program.py @@ -23,7 +23,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 @@ -39,29 +39,27 @@ __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', 'Direction', 'Divide', 'Dual', 'Expression', 'ExpressionDeclaration', - 'ExpressionNode', 'FanIn', 'FirstOf', 'Footprint', @@ -72,39 +70,40 @@ 'MaskOf', 'Multiply', 'Negate', - 'NotNode', + 'Not', 'ObjectiveDeclaration', 'ObjectiveSense', - 'OrNode', + 'Or', 'Parameter', - 'ParameterComparisonNode', + 'ParameterComparison', 'ParameterDeclaration', - 'ParameterDefinedNode', + 'ParameterDefined', 'ParameterDtype', 'Partition', 'PiecewiseDeclaration', 'Power', + 'Predicate', 'PredicateOperator', 'Program', + 'Pullback', 'QuadraticPosition', 'Reach', 'Region', - 'RelationComparisonNode', + 'RelationComparison', 'RelationDeclaration', - 'RelationDefinedNode', - 'RelationPairComparisonNode', + 'RelationDefined', + 'RelationPairComparison', 'Separability', 'SosDeclaration', 'Sum', 'Translate', - 'TypedPredicateNode', + 'TypedPredicate', 'Variable', 'VariableAbsence', 'VariableDeclaration', - 'VariableDefinedNode', + 'VariableDefined', 'VariableDomain', - 'WhereNode', - 'Window', + 'WindowSum', 'carries_variable', 'check_message', 'children', @@ -159,47 +158,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: @@ -213,30 +193,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``), @@ -244,55 +224,55 @@ 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 a relation: the dims ``direction`` consumes go, the dims it produces arrive, the dims it joins on stay. The join keys on the consumed columns and every joined column, and the operand carries every dim consumed or joined on. """ - operand: ExpressionNode + operand: Expression direction: Direction @dataclass(frozen=True) -class At(Expression): +class Pullback: """Read ``operand`` through a relation — the adjoint of :class:`GroupSum`. The dims ``direction`` consumes go and the dims it produces arrive, one value per coordinate because the read takes value columns at a key the - result fixes (``Direction.is_function_read``). The join fans out, many + result fixes, which the loader checks. The join fans out, many produced tuples sharing one consumed tuple — at each coordinate of the joined columns, which the operand carries and the result keeps. """ - operand: ExpressionNode + operand: Expression direction: Direction @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 @@ -301,17 +281,17 @@ 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 ``along`` and carries its sign in the values. - ``partition`` is a relation with a key column over ``dimension`` + ``partition`` is a relation with a key column over ``along`` (:class:`Partition`), and the translation then happens inside each group its ``within=`` columns make: the neighbour is the one before in the same group, the edge is the group's, and a wrap closes each group onto itself. A coordinate the relation sends nowhere reaches nothing. """ - operand: ExpressionNode - dimension: str + operand: Expression + along: 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 relation places nowhere reaches nothing — not even itself. """ - operand: ExpressionNode - dimension: str + operand: Expression + along: str width: int | str wrap: bool partition: Partition | 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,8 +421,9 @@ def children(expression: ExpressionNode) -> tuple[ExpressionNode, ...]: # -------------------------------------------------------------------------- -class RelationDeclaration(NamedTuple): - """One declared relation: a relation over its ``columns``, single-valued per ``key``. +@dataclass(frozen=True) +class RelationDeclaration: + """One declared relation: a table over its ``columns``, single-valued per ``key``. ``columns`` binds each role to its dimension in the order the table carries them, the key's roles first; ``key`` is the roles a row is @@ -450,7 +435,6 @@ class RelationDeclaration(NamedTuple): read one value. """ - name: str columns: tuple[tuple[str, str], ...] key: tuple[str, ...] @@ -467,30 +451,34 @@ def values(self) -> tuple[str, ...]: """The roles the key determines.""" return tuple(role for role in self.roles if role not in self.key) + @cached_property + def _dim_of(self) -> Mapping[str, str]: + """Each role's dimension, built once: :meth:`dim` is called per role inside loops over roles.""" + return dict(self.columns) + def dim(self, role: str) -> str: - return dict(self.columns)[role] + return self._dim_of[role] -class Direction(NamedTuple): +@dataclass(frozen=True) +class Direction: """One relation as one call reads it — which columns are consumed, which produced, which joined on. The declaration fixes no direction; the call does, and this is the one it - named. ``consumed``, ``produced`` and ``joined`` are *roles* — column names - of ``relation``, which binds every role to its dimension and names the key. + named. ``name`` is the relation's, as :attr:`Program.relations` keys it. + ``consumed``, ``produced`` and ``joined`` are *roles* — column names of + ``relation``, which binds every role to its dimension and names the key. ``joined`` is the key roles the call did not name (every role, for a bare relation): the join keys on them, and a value role left unnamed is not read. """ + name: str relation: RelationDeclaration consumed: tuple[str, ...] produced: tuple[str, ...] joined: tuple[str, ...] - @property - def name(self) -> str: - return self.relation.name - def dim(self, role: str) -> str: """The dimension *role* is bound to.""" return self.relation.dim(role) @@ -507,15 +495,12 @@ def produced_dims(self) -> tuple[str, ...]: def joined_dims(self) -> tuple[str, ...]: return tuple(self.dim(role) for role in self.joined) - @property - def is_function_read(self) -> bool: - """Whether the read is one value per coordinate: the key lies inside what is fixed.""" - return set(self.relation.key) <= {*self.joined, *self.produced} - -class Partition(NamedTuple): +@dataclass(frozen=True) +class Partition: """One relation as a partition steps along it — the key column stepped along, the group columns, and the key columns joined on. + ``name`` is the relation's, as :attr:`Program.relations` keys it. ``along``, ``group`` and ``joined`` are *roles* — column names of ``relation``, which binds every role to its dimension and names the key. ``along`` is the one key column over the dimension stepped along, and @@ -525,15 +510,12 @@ class Partition(NamedTuple): produced: the frame does not change. """ + name: str relation: RelationDeclaration along: str group: tuple[str, ...] joined: tuple[str, ...] - @property - def name(self) -> str: - return self.relation.name - def dim(self, role: str) -> str: """The dimension *role* is bound to.""" return self.relation.dim(role) @@ -549,9 +531,8 @@ def joined_dims(self) -> tuple[str, ...]: @dataclass(frozen=True) class DimensionDeclaration: - """A dimension and the relations with a column over it.""" + """A dimension, as the file declares it.""" - relations: tuple[RelationDeclaration, ...] = () #: 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 @@ -721,8 +702,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' @@ -737,9 +718,9 @@ class ConstraintDeclaration: """ dims: tuple[str, ...] - lhs: ExpressionNode + lhs: Expression sense: ConstraintSense - rhs: ExpressionNode + rhs: Expression where: Mask | None = None @@ -751,7 +732,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 @@ -769,7 +750,7 @@ class ObjectiveDeclaration: """Objective — scalar, every reduction in it one the file wrote.""" sense: ObjectiveSense - expression: ExpressionNode + expression: Expression @dataclass(frozen=True) @@ -778,13 +759,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 @@ -800,21 +781,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. + kinds: 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 + kinds: frozenset[type[Expression]] @dataclass(frozen=True) @@ -949,6 +922,7 @@ class Program: #: whose answer is whether the constraints can be met at all. objective: ObjectiveDeclaration | None dimensions: Mapping[str, DimensionDeclaration] = Sealed({}) + relations: Mapping[str, RelationDeclaration] = Sealed({}) sos: Mapping[str, SosDeclaration] = Sealed({}) #: Each ``piecewise:`` block the file wrote, as facts — see #: :class:`PiecewiseDeclaration`. @@ -958,7 +932,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.""" @@ -967,16 +941,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) @@ -989,23 +963,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)), + kinds=frozenset(type(node) for node in walk(*self.roots)), ) - def dimension(self, name: str) -> DimensionDeclaration: - return _declared(self.dimensions, name, 'dimension') - - @property - def relations(self) -> dict[str, RelationDeclaration]: - """Every relation in the program by name, each once — a relation keyed by two dimensions sits under both.""" - return {lk.name: lk for d in self.dimensions.values() for lk in d.relations} - - 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. @@ -1036,7 +996,7 @@ def separability(self) -> Mapping[str, Separability]: # -------------------------------------------------------------------------- -def walk_regions(*expressions: ExpressionNode) -> Iterator[tuple[ExpressionNode, tuple[Mask, ...]]]: +def walk_regions(*expressions: Expression) -> Iterator[tuple[Expression, tuple[Mask, ...]]]: """Every node under *expressions*, each with the regions it stands inside, outermost first. The traversal every *question* about a program is a filter of — which names @@ -1057,8 +1017,8 @@ def walk_regions(*expressions: ExpressionNode) -> Iterator[tuple[ExpressionNode, def _walk_regions( - expressions: tuple[ExpressionNode, ...], above: tuple[Mask, ...] -) -> Iterator[tuple[ExpressionNode, tuple[Mask, ...]]]: + expressions: tuple[Expression, ...], above: tuple[Mask, ...] +) -> Iterator[tuple[Expression, tuple[Mask, ...]]]: """The recursion under :func:`walk_regions`, with the regions above *expressions* carried down. A ``Cases`` descends by its regions rather than by :func:`children`, because @@ -1075,7 +1035,7 @@ def _walk_regions( yield from _walk_regions(children(expression), above) -def walk(*expressions: ExpressionNode) -> Iterator[ExpressionNode]: +def walk(*expressions: Expression) -> Iterator[Expression]: """Every node under *expressions*, each expression itself included, parents first. :func:`walk_regions` with the regions dropped, for the questions that do @@ -1084,7 +1044,7 @@ def walk(*expressions: ExpressionNode) -> Iterator[ExpressionNode]: return (node for node, _ in walk_regions(*expressions)) -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 — @@ -1103,22 +1063,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 @@ -1129,7 +1089,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))) @@ -1143,12 +1103,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 @@ -1161,7 +1121,7 @@ class ParameterDefinedNode: @dataclass(frozen=True) -class VariableDefinedNode: +class VariableDefined: """True at the coordinates where the named variable exists.""" name: str @@ -1169,7 +1129,7 @@ class VariableDefinedNode: @dataclass(frozen=True) -class ParameterComparisonNode: +class ParameterComparison: """Compare a parameter against a literal, element-wise.""" name: str @@ -1179,7 +1139,7 @@ class ParameterComparisonNode: @dataclass(frozen=True) -class DimensionComparisonNode: +class DimensionComparison: """Compare a dimension's own coordinates against a literal.""" name: str @@ -1188,7 +1148,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 a @@ -1203,7 +1163,7 @@ class DimensionPositionNode: @dataclass(frozen=True) -class RelationComparisonNode: +class RelationComparison: """Compare one value column of a keyed relation against a literal — ``period_of == 2030``. ``column`` is the role read, and ``dims`` the dimensions of the key @@ -1218,7 +1178,7 @@ class RelationComparisonNode: @dataclass(frozen=True) -class RelationPairComparisonNode: +class RelationPairComparison: """Compare a value column of one keyed relation with one of another — ``from_bus != to_bus`` — row by row on the key. Both keys are over the same ``dims``, and the two columns are over one @@ -1234,7 +1194,7 @@ class RelationPairComparisonNode: @dataclass(frozen=True) -class RelationDefinedNode: +class RelationDefined: """True where the relation has a row at the frame's coordinates. ``dims`` is what the frame supplies: the key's dimensions, whose row is @@ -1247,77 +1207,77 @@ class RelationDefinedNode: @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 - | RelationComparisonNode - | RelationPairComparisonNode - | RelationDefinedNode - | NotNode - | AndNode - | OrNode +Predicate = ( + BooleanLiteral + | DimensionPosition + | ParameterDefined + | VariableDefined + | ParameterComparison + | DimensionComparison + | RelationComparison + | RelationPairComparison + | RelationDefined + | 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 - | RelationComparisonNode - | RelationPairComparisonNode - | RelationDefinedNode +TypedPredicate = ( + ParameterComparison + | ParameterDefined + | VariableDefined + | DimensionComparison + | DimensionPosition + | RelationComparison + | RelationPairComparison + | RelationDefined ) #: 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. @@ -1325,9 +1285,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: @@ -1335,7 +1295,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 @@ -1348,19 +1308,19 @@ def _atom_dims(atom: TypedPredicateNode) -> frozenset[str]: 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(): + case DimensionComparison(): return frozenset({atom.name}) - case DimensionPositionNode(): + case DimensionPosition(): return frozenset({atom.name, *(atom.partition.joined_dims if atom.partition is not None else ())}) - case RelationComparisonNode() | RelationPairComparisonNode() | RelationDefinedNode(): + case RelationComparison() | RelationPairComparison() | RelationDefined(): return frozenset(atom.dims) 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 @@ -1370,23 +1330,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() - | RelationComparisonNode() - | RelationDefinedNode() - ): + case ParameterComparison() | ParameterDefined() | VariableDefined() | RelationComparison() | RelationDefined(): return frozenset({atom.name}) - case RelationPairComparisonNode(): + case RelationPairComparison(): 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 @@ -1395,12 +1349,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 @@ -1409,27 +1363,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 @@ -1446,14 +1400,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 @@ -1462,7 +1416,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) @@ -1483,12 +1437,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 56f782b3..61396994 100644 --- a/src/math_spec/resolution.py +++ b/src/math_spec/resolution.py @@ -64,25 +64,25 @@ unknown_operator_message, ) from math_spec.program import ( - AndNode, - BooleanLiteralNode, - DimensionComparisonNode, - DimensionPositionNode, + And, + BooleanLiteral, + DimensionComparison, + DimensionPosition, Direction, Mask, - NotNode, - OrNode, - ParameterComparisonNode, - ParameterDefinedNode, + Not, + Or, + ParameterComparison, + ParameterDefined, Partition, + Predicate, PredicateOperator, - RelationComparisonNode, + RelationComparison, RelationDeclaration, - RelationDefinedNode, - RelationPairComparisonNode, - TypedPredicateNode, - VariableDefinedNode, - WhereNode, + RelationDefined, + RelationPairComparison, + TypedPredicate, + VariableDefined, ) if TYPE_CHECKING: @@ -139,7 +139,7 @@ def of(cls, schema: Spec) -> Namespace: schema.variables, schema.parameters, schema.dimensions, - {n: RelationDeclaration(n, lk.pairs, lk.key_roles) for n, lk in schema.relations.items()}, + {n: RelationDeclaration(lk.pairs, lk.key_roles) for n, lk in schema.relations.items()}, { **{p: pd.dtype for p, pd in schema.parameters.items()}, **{d: dd.dtype for d, dd in schema.dimensions.items()}, @@ -276,7 +276,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. @@ -295,9 +295,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) @@ -326,22 +326,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( @@ -350,7 +350,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: @@ -670,15 +670,16 @@ def _direction( ) return None joined = tuple(r for r in shape.key if r not in from_roles and r not in into_roles) - direction = Direction(shape, from_roles, into_roles, joined) - if not forward and not direction.is_function_read: + single_valued = set(shape.key) <= {*into_roles, *joined} + direction = Direction(name, shape, from_roles, into_roles, joined) + if not forward and not single_valued: self.errors.append( f"{context}: {call}: at reads one value per coordinate, and '{name}' is not single-valued in " f'{list(from_roles)} at the columns the call lands on ({[*into_roles, *joined]}) — its key is ' f'{list(shape.key)}. Key the table by the columns the call lands on, or read the other way.' ) return None - if forward and direction.is_function_read: + if forward and single_valued: self.errors.append( f'{context}: {call}: this sum lands on the key {list(shape.key)}, so each coordinate has one ' f"term and nothing is added up — that is a read, which is at()'s. Write " @@ -741,7 +742,7 @@ def _partition( return None (along,) = over_keys joined = tuple(r for r in shape.key if r != along) - return Partition(shape, along, within_roles, joined) + return Partition(name, shape, along, within_roles, joined) def _not_a_relation(self, name: str, operator: str, key: str) -> str | None: """Why *name* is not a relation; ``None`` where it is one.""" @@ -768,27 +769,27 @@ def _not_a_relation(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) 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 relation's definedness, or a variable's existence.""" ns, context = self.ns, self.context kind = ns.kind(node.name) @@ -797,7 +798,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 " @@ -814,7 +815,7 @@ def _where_name(self, node: UnresolvedNameNode) -> WhereNode | UnresolvedWhereNo f'{node.name}.{shape.values[0] if shape.values else shape.roles[-1]} == ....' ) return node - return RelationDefinedNode(node.name, dims) + return RelationDefined(node.name, dims) case 'variable': if node.name == self.self_variable: self.errors.append( @@ -823,10 +824,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 _comparison(self, node: UnresolvedComparisonNode) -> WhereNode | UnresolvedWhereNode: + def _comparison(self, node: UnresolvedComparisonNode) -> Predicate | UnresolvedWhereNode: """``side side``, read for what each side is: a ``position()`` call, or a name against a literal or a column. The grammar admits any arithmetic on a side, and this is where the @@ -861,7 +862,7 @@ def _plain(self, node: UnresolvedComparisonNode) -> _Plain | None: def _position( self, call: FunctionCallNode, node: UnresolvedComparisonNode - ) -> DimensionPositionNode | UnresolvedComparisonNode: + ) -> DimensionPosition | UnresolvedComparisonNode: """``position(dim[, by=relation, within=columns]) i``: the name a dimension, ``by=`` a relation keyed over it.""" ns, context = self.ns, self.context shape = _position_shape(call) @@ -888,7 +889,7 @@ def _position( ) return node if by is None: - return DimensionPositionNode(dimension, node.op, position) + return DimensionPosition(dimension, node.op, position) if (problem := self._not_a_relation(by, 'position', 'by')) is not None: self.errors.append(problem) return node @@ -903,9 +904,9 @@ def _position( partition = self._partition(by, 'position', dimension, into) if partition is None: return node - return DimensionPositionNode(dimension, node.op, position, partition) + return DimensionPosition(dimension, node.op, position, partition) - def _plain_comparison(self, node: UnresolvedComparisonNode, plain: _Plain) -> WhereNode | UnresolvedWhereNode: + def _plain_comparison(self, node: UnresolvedComparisonNode, plain: _Plain) -> Predicate | UnresolvedWhereNode: """``name literal``, or the one structural form ``relation relation``.""" ns, context = self.ns, self.context value = plain.value @@ -922,7 +923,7 @@ def _plain_comparison(self, node: UnresolvedComparisonNode, plain: _Plain) -> Wh self.errors.append(refusal) return node dims = tuple(ns.relations[left_name].dim(k) for k in ns.relations[left_name].key) - return RelationPairComparisonNode(left_name, left, right_name, right, plain.op, dims) + return RelationPairComparison(left_name, left, right_name, right, plain.op, dims) self.errors.append(_declared_rhs_error(context, plain, value, rhs_kind)) return node @@ -954,15 +955,13 @@ def _plain_comparison(self, node: UnresolvedComparisonNode, plain: _Plain) -> Wh match kind: case 'parameter': assert not isinstance(value, datetime.date) - return ParameterComparisonNode(left_name, plain.op, value, ns.leaf_dims[left_name]) + return ParameterComparison(left_name, plain.op, value, ns.leaf_dims[left_name]) case 'dimension': - return DimensionComparisonNode(left_name, plain.op, value) + return DimensionComparison(left_name, plain.op, value) case 'relation': assert column is not None shape = ns.relations[left_name] - return RelationComparisonNode( - left_name, column, plain.op, value, tuple(shape.dim(k) for k in shape.key) - ) + return RelationComparison(left_name, column, plain.op, value, tuple(shape.dim(k) for k in shape.key)) case 'variable': self.errors.append( f"{context}: where references variable '{left_name}'. A where " diff --git a/src/math_spec/separability.py b/src/math_spec/separability.py index 66307bbf..2c559b28 100644 --- a/src/math_spec/separability.py +++ b/src/math_spec/separability.py @@ -9,23 +9,23 @@ from typing import TYPE_CHECKING, Literal, NamedTuple 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 class _Block(NamedTuple): @@ -39,7 +39,7 @@ class _Block(NamedTuple): label: str row: str | None - nodes: tuple[ExpressionNode, ...] + nodes: tuple[Expression, ...] mask: Mask | None reductions_couple: bool @@ -114,11 +114,11 @@ def waits_on(dimension: str, label: str, name: str, kind: Literal['offset', 'par label, f'groups {dimension} into {", ".join(node.direction.produced_dims)} — window that dimension instead, or cut only at the group edges', ) - elif isinstance(node, At): + elif isinstance(node, Pullback): for dimension in node.direction.consumed_dims: waits_on(dimension, label, node.direction.name, 'coordinate') - elif isinstance(node, (Translate, Window)): - dimension = node.dimension + elif isinstance(node, (Translate, WindowSum)): + dimension = node.along if node.wrap: report( 'coupled', @@ -130,7 +130,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.name, 'partition') - if isinstance(node, Window): + if isinstance(node, WindowSum): continue if isinstance(node.offset, str): waits_on(dimension, label, node.offset, 'offset') @@ -138,7 +138,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 d5a6684f..32a2ac2a 100644 --- a/src/math_spec/typesetting/walk.py +++ b/src/math_spec/typesetting/walk.py @@ -36,22 +36,22 @@ ) from math_spec.dimensions import dims_of from math_spec.program import ( - AndNode, - BooleanLiteralNode, - DimensionComparisonNode, - DimensionPositionNode, + And, + BooleanLiteral, + DimensionComparison, + DimensionPosition, Direction, Mask, - NotNode, - OrNode, - ParameterComparisonNode, - ParameterDefinedNode, + Not, + Or, + ParameterComparison, + ParameterDefined, + Predicate, PredicateOperator, - RelationComparisonNode, - RelationDefinedNode, - RelationPairComparisonNode, - VariableDefinedNode, - WhereNode, + RelationComparison, + RelationDefined, + RelationPairComparison, + VariableDefined, ) from math_spec.typesetting.format import Entry, Glossary, Line, OperatorName @@ -557,33 +557,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 ( @@ -591,7 +591,7 @@ def _where(self, node: WhereNode, ctx: _Context) -> tuple[str, int]: comparison, ) - if isinstance(node, DimensionPositionNode): + if isinstance(node, DimensionPosition): grouping = ( None if node.partition is None @@ -601,30 +601,30 @@ def _where(self, node: WhereNode, ctx: _Context) -> tuple[str, int]: ordinal = self._ordinal(node.name, node.position, grouping) return f'{place} {self._op(_PREDICATES[node.op])} {ordinal}', comparison - if isinstance(node, RelationComparisonNode): + if isinstance(node, RelationComparison): applied = self._value_read(node.name, node.column, ctx) return f'{applied} {self._op(_PREDICATES[node.op])} {self._literal(node.value)}', comparison - if isinstance(node, RelationPairComparisonNode): + if isinstance(node, RelationPairComparison): left = self._value_read(node.name, node.column, ctx) right = self._value_read(node.other, node.other_column, ctx) return f'{left} {self._op(_PREDICATES[node.op])} {right}', comparison - if isinstance(node, RelationDefinedNode): + if isinstance(node, RelationDefined): return self._relation_row(node.name, self._frame_key(node.name, ctx)), 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 e292840e..f864d790 100644 --- a/src/math_spec/validation.py +++ b/src/math_spec/validation.py @@ -38,7 +38,7 @@ from math_spec.expansion import expand, parse_and_expand, parse_template from math_spec.model import Spec from math_spec.operators import BUILTINS, call_shape_error, unknown_operator_message -from math_spec.program import BooleanLiteralNode +from math_spec.program import BooleanLiteral from math_spec.resolution import ( Namespace, Resolved, @@ -53,7 +53,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 | Mapping[str, object] | Spec) -> Spec: @@ -186,11 +186,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 5258042b..ccfe4e29 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 Mask, RelationPairComparisonNode +from math_spec.program import Mask, RelationPairComparison from math_spec.resolution import Namespace, expression_of, where_of from math_spec.validation import to_spec from tests.fixtures import override, schema_of @@ -533,6 +533,6 @@ def test_names_read_takes_both_sides_of_a_relation_pair(): BASE has one relation per dimension, so the pair is built directly rather than resolved from a predicate string. """ - where = RelationPairComparisonNode('from_bus', 'bus', 'to_bus', 'bus', '!=', ('line',)) + where = RelationPairComparison('from_bus', 'bus', 'to_bus', 'bus', '!=', ('line',)) assert Mask(where).names_read == {'from_bus', 'to_bus'}, 'a relation pair names both maps it compares' diff --git a/tests/test_exclusivity.py b/tests/test_exclusivity.py index 4541d2b5..c6b9b7e2 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 6420f866..e198898b 100644 --- a/tests/test_lowering.py +++ b/tests/test_lowering.py @@ -24,36 +24,36 @@ from math_spec.program import ( QUADRATIC_POSITIONS, Add, - AndNode, - At, - BooleanLiteralNode, + And, + BooleanLiteral, Cases, Constant, - DimensionComparisonNode, + DimensionComparison, DimensionDeclaration, Direction, Divide, Dual, - ExpressionNode, + Expression, Footprint, GroupSum, Mask, Multiply, Negate, - NotNode, - OrNode, + Not, + Or, Parameter, - ParameterComparisonNode, - ParameterDefinedNode, + ParameterComparison, + ParameterDefined, Partition, Power, Program, + Pullback, Region, RelationDeclaration, Sum, Translate, Variable, - Window, + WindowSum, children, divisor_parameters, fan_in, @@ -72,7 +72,7 @@ DISPATCH_YAML = EXAMPLES / 'dispatch.yaml' #: The mask `examples/dispatch.yaml` puts on `dispatch`, as the plan carries it. -CAPACITY_POSITIVE = ParameterComparisonNode('capacity', '>', 0.0, ('generator',)) +CAPACITY_POSITIVE = ParameterComparison('capacity', '>', 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 @@ -85,10 +85,10 @@ } #: `lk` as `sum` reads it: key consumed, value produced, nothing joined. -LK = RelationDeclaration('lk', (('g', 'g'), ('h', 'h')), ('g',)) -LK2 = RelationDeclaration('lk2', (('g', 'g'), ('z', 'z')), ('g',)) -LK_DIRECTION = Direction(LK, ('g',), ('h',), ()) -AT_BUS = RelationDeclaration('at_bus', (('g', 'g'), ('bus', 'bus')), ('g',)) +LK = RelationDeclaration((('g', 'g'), ('h', 'h')), ('g',)) +LK2 = RelationDeclaration((('g', 'g'), ('z', 'z')), ('g',)) +LK_DIRECTION = Direction('lk', LK, ('g',), ('h',), ()) +AT_BUS = RelationDeclaration((('g', 'g'), ('bus', 'bus')), ('g',)) #: `fixtures.SMALL_MODEL` plus a second relation and a per-entity #: offset. Which node a construct becomes is mostly a claim about the dim it @@ -151,7 +151,7 @@ def test_lower_program_structure(dispatch_program): assert dispatch_program.objective.sense == 'minimize', "the program carries the language's spelling, untranslated" assert dispatch_program.objective.expression == Sum( - Variable('dispatch') * Parameter('cost'), ('generator', 'snapshot') + Multiply(Variable('dispatch'), Parameter('cost')), ('generator', 'snapshot') ), 'the objective carries the sum the file wrote, over the dims it named none of' @@ -182,33 +182,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('capacity', ParameterDefinedNode('capacity', ('generator',)), id='a-bare-parameter-name'), + pytest.param('capacity', ParameterDefined('capacity', ('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( 'capacity > 0 AND NOT load == 0', - AndNode(CAPACITY_POSITIVE, NotNode(ParameterComparisonNode('load', '==', 0.0, ('snapshot',)))), + And(CAPACITY_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('capacity > 0 AND True', CAPACITY_POSITIVE, id='and-true-is-the-other-side'), pytest.param('capacity > 0 OR False', CAPACITY_POSITIVE, id='or-false-is-the-other-side'), pytest.param('capacity > 0 OR True', None, id='or-true-is-no-mask-at-all'), - pytest.param('capacity > 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('capacity > 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 (capacity > 0 AND False)', None, id='a-branch-folded-away-folds-the-one-above-it'), pytest.param( 'NOT (NOT capacity)', - ParameterDefinedNode('capacity', ('generator',)), + ParameterDefined('capacity', ('generator',)), id='a-double-negation-cancels-on-the-load-path', ), pytest.param( '(capacity > 0 OR True) AND load', - ParameterDefinedNode('load', ('snapshot',)), + ParameterDefined('load', ('snapshot',)), id='an-absorbed-side-takes-its-own-branch-with-it', ), ], @@ -216,7 +216,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, ( @@ -230,7 +230,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): @@ -306,17 +306,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(CAPACITY_POSITIVE), (CAPACITY_POSITIVE,), id='a-not-carries-its-operand'), - pytest.param(AndNode(CAPACITY_POSITIVE, FLAG), (CAPACITY_POSITIVE, FLAG), id='an-and-carries-both-sides'), - pytest.param(OrNode(CAPACITY_POSITIVE, FLAG), (CAPACITY_POSITIVE, FLAG), id='an-or-carries-both-sides'), + pytest.param(Not(CAPACITY_POSITIVE), (CAPACITY_POSITIVE,), id='a-not-carries-its-operand'), + pytest.param(And(CAPACITY_POSITIVE, FLAG), (CAPACITY_POSITIVE, FLAG), id='an-and-carries-both-sides'), + pytest.param(Or(CAPACITY_POSITIVE, FLAG), (CAPACITY_POSITIVE, FLAG), id='an-or-carries-both-sides'), pytest.param(CAPACITY_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): @@ -334,32 +334,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(CAPACITY_POSITIVE)).dims == {'generator'}, 'negation keeps the dims it negates' + assert Mask(Not(CAPACITY_POSITIVE)).dims == {'generator'}, 'negation keeps the dims it negates' assert (Mask(CAPACITY_POSITIVE) & Mask(b)).dims == {'generator', 'snapshot'}, 'conjunction unions both sides' - assert (Mask(CAPACITY_POSITIVE) & Mask(b)).root == AndNode(CAPACITY_POSITIVE, b), ( + assert (Mask(CAPACITY_POSITIVE) & Mask(b)).root == And(CAPACITY_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' @@ -371,9 +371,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' @@ -398,7 +398,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): @@ -418,7 +418,7 @@ def test_a_power_lowers_to_a_node_of_its_own(dispatch_schema): ), pytest.param( 'at(r, by=lk, over=h, into=g)', - At(Variable('r'), direction=Direction(LK, ('h',), ('g',), ())), + Pullback(Variable('r'), direction=Direction('lk', LK, ('h',), ('g',), ())), id='a-pullback-reads-the-same-table-back', ), pytest.param( @@ -444,28 +444,28 @@ def test_a_power_lowers_to_a_node_of_its_own(dispatch_schema): offset=1, wrap=False, fill=0.0, - partition=Partition(LK, 'g', ('h',), ()), + partition=Partition('lk', LK, 'g', ('h',), ()), ), id='a-translation-stops-at-the-edges-of-the-relation-it-names', ), pytest.param( 'sum_back(p, along=g, window=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, along=g, window=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, along=g, window=2, by=lk, within=h)', - Window( + WindowSum( Variable('p'), 'g', width=2, wrap=False, - partition=Partition(LK, 'g', ('h',), ()), + partition=Partition('lk', LK, 'g', ('h',), ()), ), id='a-window-stops-at-the-edges-of-the-relation-it-names', ), @@ -508,7 +508,7 @@ def test_a_partition_keeps_its_group_when_the_relation_gains_a_value_column(): def _partition_of(row): """The one partition a constraint row's expression carries.""" nodes = [*walk(row.lhs), *walk(row.rhs)] - [partition] = [node.partition for node in nodes if isinstance(node, Translate | Window)] + [partition] = [node.partition for node in nodes if isinstance(node, Translate | WindowSum)] return partition @@ -544,14 +544,12 @@ def test_a_relation_lowers_with_the_direction_each_call_names(): ) columns = (('generator', 'generator'), ('snapshot', 'snapshot'), ('zone', 'zone')) - declared = RelationDeclaration('zone_of', columns, ('generator', 'snapshot')) - assert program.dimension('generator').relations == (declared,), 'the relation sits under its first column' - assert program.dimension('zone').relations == (declared,), 'and under its last' - assert program.relations == {'zone_of': declared}, 'and once in the program' + declared = RelationDeclaration(columns, ('generator', 'snapshot')) + assert program.relations == {'zone_of': declared}, 'the relation sits once in the program, under its name' zonal = program.constraints['zonal'].lhs - assert zonal == GroupSum(Variable('p'), direction=Direction(declared, ('generator',), ('zone',), ('snapshot',))), ( - 'a grouped sum names the column it consumes, the one it produces and the one it joins on' - ) + assert zonal == GroupSum( + Variable('p'), direction=Direction('zone_of', declared, ('generator',), ('zone',), ('snapshot',)) + ), 'a grouped sum names the column it consumes, the one it produces and the one it joins on' assert isinstance(zonal, GroupSum) assert (zonal.direction.consumed_dims, zonal.direction.produced_dims, zonal.direction.joined_dims) == ( ('generator',), @@ -562,25 +560,25 @@ def test_a_relation_lowers_with_the_direction_each_call_names(): 'the direction holds the one declaration the program holds, not an equal copy built again' ) assert program.constraints['history'].lhs == GroupSum( - Variable('p'), direction=Direction(declared, ('snapshot',), ('zone',), ('generator',)) + Variable('p'), direction=Direction('zone_of', declared, ('snapshot',), ('zone',), ('generator',)) ), 'the same table read from its other key column' priced = program.constraints['priced'].rhs - assert priced == At(Parameter('price'), direction=Direction(declared, ('zone',), ('generator',), ('snapshot',))), ( - 'and its adjoint consumes the value column and produces the key column' - ) - assert isinstance(priced, At) + assert priced == Pullback( + Parameter('price'), direction=Direction('zone_of', declared, ('zone',), ('generator',), ('snapshot',)) + ), 'and its adjoint consumes the value column and produces the key column' + assert isinstance(priced, Pullback) assert (priced.direction.consumed_dims, priced.direction.produced_dims, priced.direction.joined_dims) == ( ('zone',), ('generator',), ('snapshot',), ), 'an at consumes the coarse dims, produces the fine, and joins on the rest of the key' - p_where = program.variable('p').where + p_where = program.variables['p'].where assert p_where is not None assert [(type(a).__name__, a.dims) for a in p_where.atoms] == [ - ('RelationComparisonNode', ('generator', 'snapshot')), - ('RelationDefinedNode', ('generator', 'snapshot')), + ('RelationComparison', ('generator', 'snapshot')), + ('RelationDefined', ('generator', 'snapshot')), ], 'a comparison and an existence are both read at the key of a keyed relation' - first_where = program.variable('first').where + first_where = program.variables['first'].where assert first_where is not None assert first_where.dims == {'generator', 'snapshot'}, 'a position within a group is read at every key column' @@ -589,16 +587,16 @@ def test_a_binary_variable_lowers_to_a_binary_domain(): program = to_program( schema_of(DISPATCH_YAML, **{'variables.dispatch.domain': 'binary', 'variables.dispatch.bounds': {}}) ) - assert program.variable('dispatch').domain == 'binary' + assert program.variables['dispatch'].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')) - component_of = RelationDeclaration('component_of', (('flow', 'flow'), ('component', 'component')), ('flow',)) - pulled = At(quotient, direction=Direction(component_of, ('component',), ('flow',), ())) + component_of = RelationDeclaration((('flow', 'flow'), ('component', 'component')), ('flow',)) + pulled = Pullback(quotient, direction=Direction('component_of', component_of, ('component',), ('flow',), ())) - 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' @@ -616,18 +614,18 @@ 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' ) -OUTER = Mask(ParameterDefinedNode('committable', ('g',))) -INNER = Mask(ParameterDefinedNode('flag', ('g',))) +OUTER = Mask(ParameterDefined('committable', ('g',))) +INNER = Mask(ParameterDefined('flag', ('g',))) NESTED = Add( Variable('x'), Cases( @@ -671,11 +669,11 @@ def test_walk_is_the_node_column_of_walk_regions(): Power(Parameter('c'), Constant(2.0)): 'one-to-one', Divide(Variable('p'), Parameter('c')): 'one-to-one', Sum(Variable('p'), ('g',)): 'many-to-one', - GroupSum(Variable('p'), direction=Direction(AT_BUS, ('g',), ('bus',), ())): 'many-to-one', - At(Variable('p'), direction=Direction(AT_BUS, ('bus',), ('g',), ())): 'one-to-one', + GroupSum(Variable('p'), direction=Direction('at_bus', AT_BUS, ('g',), ('bus',), ())): 'many-to-one', + Pullback(Variable('p'), direction=Direction('at_bus', AT_BUS, ('bus',), ('g',), ())): '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', } @@ -683,8 +681,8 @@ def test_walk_is_the_node_column_of_walk_regions(): 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' ) @@ -693,38 +691,25 @@ def test_a_node_answers_its_fan_in(node, expected): assert fan_in(node) == expected -def test_a_relation_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( - (RelationDeclaration('season_of', (('snapshot', 'snapshot'), ('season', 'season')), ('snapshot',)),) - ), - 'generator': DimensionDeclaration( - (RelationDeclaration('at_bus', (('generator', 'generator'), ('bus', 'bus')), ('generator',)),) - ), - }, - ) - - assert [lk.name for lk in program.dimension('snapshot').relations] == ['season_of'], ( - 'one dimension names its own maps and no other dimension' +def test_a_relation_is_declared_as_the_file_declares_it(): + """One group keyed by name, each entry its columns and its key, and nothing nested under a dimension.""" + program = to_program( + override( + TINY, + dimensions={'g': {}, 'bus': {}, 'season': {}}, + relations={ + 'season_of': {'key': 'g', 'values': 'season'}, + 'at_bus': {'key': 'g', 'values': 'bus'}, + }, + ) ) - assert program.dimension('snapshot').relations[0].values == ('season',), 'and the map says what its key determines' - assert list(program.relations) == ['season_of', 'at_bus'], 'every map once, by name, 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.relations == { + 'season_of': RelationDeclaration((('g', 'g'), ('season', 'season')), ('g',)), + 'at_bus': RelationDeclaration((('g', 'g'), ('bus', 'bus')), ('g',)), + }, 'every relation under its own name, in declaration order' + assert program.relations['season_of'].values == ('season',), 'and each says what its key determines' + assert program.dimensions['g'] == DimensionDeclaration(dtype='str'), 'a dimension carries its dtype and no relation' def test_a_program_is_built_by_keyword_so_a_field_added_later_cannot_reorder_an_old_call(): @@ -733,15 +718,15 @@ def test_a_program_is_built_by_keyword_so_a_field_added_later_cannot_reorder_an_ Program({}, {}, {}, None) # pyrefly: ignore[bad-argument-count] the point of the test -@pytest.mark.parametrize('group', ['parameters', 'variables', 'constraints', 'dimensions', 'sos']) +@pytest.mark.parametrize('group', ['parameters', 'variables', 'constraints', 'dimensions', 'relations', 'sos']) def test_a_program_seals_its_declaration_groups(dispatch_program, group): """`frozen=True` sealed the fields and said nothing about what was behind them.""" with pytest.raises(TypeError): 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, @@ -750,16 +735,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, ( - 'a named expression builds no row, so it is not one of the expressions a row is built from' + assert program.expressions['spend'].expression not in program.roots, ( + 'a named expression builds no row, so it is not one of the trees 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: @@ -798,7 +783,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.kinds)} == {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' @@ -816,8 +801,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.kinds, "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' ) @@ -830,8 +815,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 = { @@ -878,10 +863,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' ) @@ -921,7 +904,7 @@ 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 = 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' ) @@ -955,7 +938,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(): @@ -967,7 +950,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' ) @@ -983,7 +966,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' ) @@ -991,7 +974,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 afce6c51..d0d73806 100644 --- a/tests/test_parser.py +++ b/tests/test_parser.py @@ -45,11 +45,11 @@ ) from math_spec.errors import SchemaError from math_spec.program import ( - AndNode, - BooleanLiteralNode, + And, + BooleanLiteral, Direction, - NotNode, - OrNode, + Not, + Or, Partition, RelationDeclaration, _conjuncts, @@ -61,13 +61,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( @@ -262,12 +262,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': '>', 'right': NumberNode(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): @@ -280,8 +280,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')) ) @@ -295,7 +295,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.""" @@ -390,7 +390,7 @@ def test_a_side_is_any_arithmetic_to_the_grammar(text): def test_a_bracketed_predicate_is_still_a_predicate(): """`(a > 0) AND b` groups a comparison; only `(a + b) <= c` brackets arithmetic.""" - assert isinstance(parse_where('(a > 0) AND b'), AndNode) + assert isinstance(parse_where('(a > 0) AND b'), And) def test_a_where_side_is_held_to_the_depth_an_expression_is(): @@ -543,7 +543,7 @@ def test_a_node_prints_as_the_file_writes_it(text, printed): assert str(parse_expression(text)) == printed, 'the spelling is the one a file could be written with' -_ZONE_OF = RelationDeclaration('zone_of', (('u', 'unit'), ('zone', 'zone')), ('u',)) +_ZONE_OF = RelationDeclaration((('u', 'unit'), ('zone', 'zone')), ('u',)) @pytest.mark.parametrize( @@ -554,12 +554,12 @@ def test_a_node_prints_as_the_file_writes_it(text, printed): pytest.param(DimensionNode('t'), 't', id='a-dimension'), pytest.param(DualNode('budget'), 'dual(budget)', id='a-dual'), pytest.param( - DirectionNode(Direction(_ZONE_OF, ('u',), ('zone',), ())), + DirectionNode(Direction('zone_of', _ZONE_OF, ('u',), ('zone',), ())), 'zone_of', id='a-relation-read-in-a-direction', ), pytest.param( - PartitionNode(Partition(_ZONE_OF, 'u', ('zone',), ())), + PartitionNode(Partition('zone_of', _ZONE_OF, 'u', ('zone',), ())), 'zone_of', id='a-relation-stepped-along-as-a-partition', ), diff --git a/tests/test_piecewise.py b/tests/test_piecewise.py index 48d71537..d951b26a 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 4e8fefb3..3b31c8d5 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' @@ -525,7 +525,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.partition.name if node.partition is not None else None) == by diff --git a/tests/typesetting/test_golden.py b/tests/typesetting/test_golden.py index 45d74650..1966268e 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 @@ -148,10 +148,11 @@ def _rendered_trees() -> Iterator[object]: } #: A dataclass the walk steps *through* rather than renders: an arm has no -#: branch of its own — its ``when`` and ``value`` do. Not a member of any node -#: union, so it is subtracted from what the tree walk finds rather than added -#: to what the vocabulary declares. -CARRIERS = {'CaseArm'} +#: branch of its own — its ``when`` and ``value`` do — and a direction and the +#: relation it reads are the facts a node carries rather than nodes. None is a +#: member of any node union, so they are subtracted from what the tree walk +#: finds rather than added to what the vocabulary declares. +CARRIERS = {'CaseArm', 'Direction', 'Partition', 'RelationDeclaration'} def test_the_golden_model_carries_every_node_kind_the_walk_renders(): @@ -163,7 +164,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, '