diff --git a/docs/contributing.md b/docs/contributing.md index 89815cdd..362126d3 100644 --- a/docs/contributing.md +++ b/docs/contributing.md @@ -126,11 +126,11 @@ stale anchor fails it. The same construct passes through three layers, and each names it in full. The suffix says which layer, which keeps the three vocabularies from colliding: -| Layer | Suffix | Example | -| ------------------------------- | -------------------- | ----------------------------------------- | -| YAML block (`math_spec.model`) | `Block` | `VariableBlock`, `PiecewiseBlock` | -| Core AST (`math_spec.*_parser`) | `Node` | `VariableNode`, `DimensionComparisonNode` | -| Program (`math_spec.program`) | none / `Declaration` | `Variable`, `VariableDeclaration` | +| Layer | Suffix | Example | +| ------------------------------- | -------------------- | ------------------------------------------ | +| YAML block (`math_spec.model`) | `Block` | `VariableBlock`, `PiecewiseBlock` | +| Core AST (`math_spec.*_parser`) | `Node` | `VariableNode`, `UnresolvedComparisonNode` | +| Program (`math_spec.program`) | none / `Declaration` | `Variable`, `VariableDeclaration` | Two rules follow, and a PR that adds a construct keeps them: diff --git a/docs/reference/language/reading.md b/docs/reference/language/reading.md index d546b5ff..117151d3 100644 --- a/docs/reference/language/reading.md +++ b/docs/reference/language/reading.md @@ -128,7 +128,7 @@ footprint = program.footprint sorted(footprint.quadratic) # [] sorted(footprint.domains) # ['continuous'] sorted(footprint.sos_types) # [] -sorted(kind.__name__ for kind in footprint.shapes) # ['Constant', 'Multiply', 'Parameter', 'Sum', 'Variable'] +sorted(kind.__name__ for kind in footprint.nodes) # ['Constant', 'Multiply', 'Parameter', 'Sum', 'Variable'] ``` Every field is a set. `if footprint.sos_types` asks whether sets appear at all, @@ -143,7 +143,7 @@ model does not use the construct, not that the construct does not exist. on the numbers. The footprint stops at the kind of construct. An engine whose solver accepts a -window but not a wrapped one reads `Window in footprint.shapes`, then walks the +window but not a wrapped one reads `WindowSum in footprint.nodes`, then walks the tree for the detail. ## Asking whether an axis can be cut diff --git a/docs/reference/language/reported.md b/docs/reference/language/reported.md index 6d6fd1a9..178befb4 100644 --- a/docs/reference/language/reported.md +++ b/docs/reference/language/reported.md @@ -44,7 +44,7 @@ variable in it, such as `(1 + rate) ** period`, is reported all the same. Deciding by use costs one thing: an entry meant for a constraint, and never named there, loads as a reported quantity instead of failing. -An engine reads the answer at `Program.named_expressions[name].in_math`. +An engine reads the answer at `Program.expressions[name].in_math`. ## Which restrictions do not apply diff --git a/src/math_spec/_expression_parser.py b/src/math_spec/_expression_parser.py index 132a5182..fb2fbcb5 100644 --- a/src/math_spec/_expression_parser.py +++ b/src/math_spec/_expression_parser.py @@ -22,7 +22,7 @@ if TYPE_CHECKING: from collections.abc import Callable, Iterator, Mapping - from math_spec.program import Walk, WhereNode + from math_spec.program import Predicate, Walk #: The relation a comparison may carry — the three an expression may be #: written with, which is what a constraint's sense is read off. @@ -191,7 +191,7 @@ class CaseArm: """ label: str - when: WhereNode | None + when: Predicate | None value: ArithmeticNode @@ -263,7 +263,7 @@ class ComparisonNode: #: A whole spec-side expression tree — parse output and the resolved tree alike. -#: Named apart from :data:`math_spec.program.ExpressionNode`, the lowered +#: Named apart from :data:`math_spec.program.Expression`, the lowered #: vocabulary a consumer reads. ParsedNode = ArithmeticNode | ComparisonNode diff --git a/src/math_spec/_where_parser.py b/src/math_spec/_where_parser.py index c0819cad..a0f18ff2 100644 --- a/src/math_spec/_where_parser.py +++ b/src/math_spec/_where_parser.py @@ -16,12 +16,12 @@ import pyparsing as pp from math_spec._expression_parser import NAME, REAL, parse_text -from math_spec.program import AndNode, BooleanLiteralNode, NotNode, OrNode, PredicateOperator, where_children +from math_spec.program import And, BooleanLiteral, Not, Or, PredicateOperator, where_children if TYPE_CHECKING: from collections.abc import Callable - from math_spec.program import WhereNode + from math_spec.program import Predicate # --------------------------------------------------------------------------- # AST nodes @@ -100,8 +100,8 @@ def _build_where_grammar() -> pp.ParserElement: """ where_expr = pp.Forward() - true_lit = pp.CaselessKeyword('True').set_parse_action(lambda: BooleanLiteralNode(True)) - false_lit = pp.CaselessKeyword('False').set_parse_action(lambda: BooleanLiteralNode(False)) + true_lit = pp.CaselessKeyword('True').set_parse_action(lambda: BooleanLiteral(True)) + false_lit = pp.CaselessKeyword('False').set_parse_action(lambda: BooleanLiteral(False)) # pyrefly: ignore[implicit-any-lambda] number = pp.Regex(rf'-?({REAL}|\d+)').set_parse_action(lambda t: float(t[0])) @@ -142,28 +142,28 @@ def _build_where_grammar() -> pp.ParserElement: NOT = pp.CaselessKeyword('NOT').suppress() # pyrefly: ignore[implicit-any-lambda] - not_expr = (NOT + atom).set_parse_action(lambda t: NotNode(t[0])) | atom + not_expr = (NOT + atom).set_parse_action(lambda t: Not(t[0])) | atom AND = pp.CaselessKeyword('AND').suppress() and_expr = not_expr + pp.ZeroOrMore(AND + not_expr) - and_expr.set_parse_action(_folder(AndNode)) + and_expr.set_parse_action(_folder(And)) OR = pp.CaselessKeyword('OR').suppress() or_expr = and_expr + pp.ZeroOrMore(OR + and_expr) - or_expr.set_parse_action(_folder(OrNode)) + or_expr.set_parse_action(_folder(Or)) where_expr <<= or_expr return where_expr -def _folder(node_type: type[AndNode] | type[OrNode]) -> Callable[[pp.ParseResults], Any]: +def _folder(node_type: type[And] | type[Or]) -> Callable[[pp.ParseResults], Any]: """A parse action left-folding a flat operator chain into *node_type*.""" def fold(tokens: pp.ParseResults) -> Any: items = list(tokens) - result: WhereNode | UnresolvedWhereNode = items[0] + result: Predicate | UnresolvedWhereNode = items[0] for item in items[1:]: - result = node_type(cast('WhereNode', result), item) + result = node_type(cast('Predicate', result), item) return result return fold @@ -200,7 +200,7 @@ def _named_rewrite(text: str, loc: int) -> str | None: @lru_cache(maxsize=4096) -def parse_where(text: str) -> WhereNode | UnresolvedWhereNode: +def parse_where(text: str) -> Predicate | UnresolvedWhereNode: """Parse a where string into an AST, its leaves still unresolved. The connectives and literals are the resolved vocabulary's own; the leaves @@ -214,6 +214,6 @@ def parse_where(text: str) -> WhereNode | UnresolvedWhereNode: complaint. """ return cast( - 'WhereNode | UnresolvedWhereNode', + 'Predicate | UnresolvedWhereNode', parse_text(_WHERE_GRAMMAR, text, 'where string', _named_rewrite, where_children, _DEEP_REWRITE), ) diff --git a/src/math_spec/advice.py b/src/math_spec/advice.py index 910478db..45e50486 100644 --- a/src/math_spec/advice.py +++ b/src/math_spec/advice.py @@ -14,7 +14,7 @@ from math_spec.boundedness import unbounded_notes from math_spec.errors import Advice from math_spec.lowering import to_program -from math_spec.program import At, GroupSum, walk +from math_spec.program import GroupSum, Pullback, walk if TYPE_CHECKING: from pathlib import Path @@ -72,9 +72,9 @@ def _produced_axes(program: Program) -> set[str]: ``sum(by=)`` lands on its target and ``at()`` spreads onto its fine dimension. """ axes: set[str] = set() - for node in walk(*program.expressions): + for node in walk(*program.roots): if isinstance(node, GroupSum): axes.update(node.into) - elif isinstance(node, At): + elif isinstance(node, Pullback): axes.update(node.over) return axes diff --git a/src/math_spec/boundedness.py b/src/math_spec/boundedness.py index 0c6b3e7e..1102abe8 100644 --- a/src/math_spec/boundedness.py +++ b/src/math_spec/boundedness.py @@ -19,21 +19,21 @@ from math_spec.errors import Advice from math_spec.program import ( Add, - At, Cases, Constant, Divide, Dual, - ExpressionNode, + Expression, GroupSum, Multiply, Negate, Parameter, Power, + Pullback, Sum, Translate, Variable, - Window, + WindowSum, children, variables_of, ) @@ -113,7 +113,7 @@ def _times(sign: Sign, other: Sign) -> Sign: return None if sign is None or other is None else ('+' if sign == other else '-') -def _coefficient_sign(node: ExpressionNode) -> Sign: +def _coefficient_sign(node: Expression) -> Sign: """The sign *node* scales a term by, or ``None`` unless it is a signed constant. ``-2`` lowers to a negation over a constant, so the sign of a literal @@ -128,7 +128,7 @@ def _coefficient_sign(node: ExpressionNode) -> Sign: return None -def _record_signs(node: ExpressionNode, sign: Sign, signs: dict[str, Sign]) -> None: +def _record_signs(node: Expression, sign: Sign, signs: dict[str, Sign]) -> None: """Record the sign each variable under *node* carries into the objective. A variable reached twice with different signs, or once with an undecidable @@ -162,7 +162,7 @@ def _record_signs(node: ExpressionNode, sign: Sign, signs: dict[str, Sign]) -> N _record_signs(node.base, None, signs) _record_signs(node.exponent, None, signs) return - if isinstance(node, Sum | GroupSum | At | Translate | Window | Cases): + if isinstance(node, Sum | GroupSum | Pullback | Translate | WindowSum | Cases): for child in children(node): _record_signs(child, sign, signs) return diff --git a/src/math_spec/dimensions.py b/src/math_spec/dimensions.py index dbb5e8a8..8b19998e 100644 --- a/src/math_spec/dimensions.py +++ b/src/math_spec/dimensions.py @@ -40,15 +40,15 @@ from math_spec.errors import DimensionError from math_spec.operators import BUILTINS from math_spec.program import ( - DimensionComparisonNode, - DimensionPositionNode, + DimensionComparison, + DimensionPosition, Mask, - ParameterComparisonNode, - ParameterDefinedNode, - RelationComparisonNode, - RelationDefinedNode, - RelationPairComparisonNode, - VariableDefinedNode, + ParameterComparison, + ParameterDefined, + RelationComparison, + RelationDefined, + RelationPairComparison, + VariableDefined, ) if TYPE_CHECKING: @@ -520,13 +520,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 75ddcb45..127d1d75 100644 --- a/src/math_spec/exclusivity.py +++ b/src/math_spec/exclusivity.py @@ -22,27 +22,27 @@ from typing import TYPE_CHECKING, Any, Literal, assert_never, cast from math_spec.program import ( - AndNode, - BooleanLiteralNode, - DimensionComparisonNode, - DimensionPositionNode, + 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) @@ -185,15 +185,15 @@ def witness(self, cell: dict[Subject, Cell]) -> str: return ', '.join(f'{subject} is {_shown(subject, value)}' for subject, value in cell.items()) -def _observe(node: TypedPredicateNode, subject: Subject, values: set[Any], dtypes: Mapping[str, DeclaredDtype]) -> None: +def _observe(node: TypedPredicate, subject: Subject, values: set[Any], dtypes: Mapping[str, DeclaredDtype]) -> None: """Record what *node* says about its subject: a position, or a literal. ``position()`` converts the dimension to an integer, so an ordering over a rank is an ordering of integers and every comparator is admitted there. """ - if isinstance(node, DimensionPositionNode): + if isinstance(node, DimensionPosition): values.add(node.position) - elif isinstance(node, 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 ' @@ -201,7 +201,7 @@ def _observe(node: TypedPredicateNode, subject: Subject, values: set[Any], dtype f'ordering as a boolean parameter and test that' ) raise Undecidable(msg) - elif isinstance(node, ParameterComparisonNode | DimensionComparisonNode | 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 ' @@ -212,21 +212,21 @@ def _observe(node: TypedPredicateNode, subject: Subject, values: set[Any], dtype values.add(node.value) -def _subject_of(node: TypedPredicateNode) -> Subject: +def _subject_of(node: TypedPredicate) -> Subject: match node: - case ParameterDefinedNode(name=name) | ParameterComparisonNode(name=name): + case ParameterDefined(name=name) | ParameterComparison(name=name): return Subject('param', name) - case VariableDefinedNode(name=name): + case VariableDefined(name=name): return Subject('variable', name) - case DimensionComparisonNode(name=name): + case DimensionComparison(name=name): return Subject('dim', name) - case DimensionPositionNode(name=name, partition=partition): + case DimensionPosition(name=name, partition=partition): if partition is None: return Subject('rank', name) return Subject('rank', name, partition.name, partition.produced) - 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) @@ -380,42 +380,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 a78fab20..9d102e5a 100644 --- a/src/math_spec/lowering.py +++ b/src/math_spec/lowering.py @@ -140,17 +140,8 @@ def lower_program(expanded: _ExpandedSpec) -> program.Program: _Lowering(expanded, 'the objective').expr(resolved.objective), ) - dimensions = { - dname: program.DimensionDeclaration( - tuple( - program.RelationDeclaration(lname, lk.pairs, lk.keys) - for lname, lk in expanded.relations.items() - 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()} + relations = {name: program.RelationDeclaration(lk.pairs, lk.keys) for name, lk in expanded.relations.items()} sos = { sname: program.SosDeclaration( sdef.variable, @@ -171,9 +162,10 @@ def lower_program(expanded: _ExpandedSpec) -> program.Program: constraints=constraints, objective=objective, dimensions=dimensions, + relations=relations, sos=sos, piecewise={name: declaration_of(ex) for name, ex in expanded.expanded_piecewise.items()}, - named_expressions=expressions, + expressions=expressions, ) @@ -189,7 +181,7 @@ class _Lowering: schema: _ExpandedSpec context: str - def expr(self, node: ArithmeticNode) -> program.ExpressionNode: + def expr(self, node: ArithmeticNode) -> program.Expression: """Rewrite one resolved core-AST expression as a program expression.""" if isinstance(node, NumberNode): return program.Constant(node.value) @@ -258,7 +250,7 @@ def _cases(self, node: CasesNode) -> program.Cases: regions.append(program.Region(when, self.expr(arm.value))) return program.Cases(tuple(regions)) - def sum(self, node: FunctionCallNode) -> program.ExpressionNode: + def sum(self, node: FunctionCallNode) -> program.Expression: """``sum(x)``, ``sum(x, over=d)`` or ``sum(x, by=relation)``. Two program nodes under one surface verb: reducing a dim away and reducing it @@ -276,13 +268,13 @@ def sum(self, node: FunctionCallNode) -> program.ExpressionNode: assert isinstance(by_node, RelationNode), 'resolution refuses a by= that is not a relation' return program.GroupSum(operand, walks=by_node.walks) - 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, RelationNode), 'resolution refuses a by= that is not a relation' - return program.At(self.expr(node.args[0]), walks=by_node.walks) + return program.Pullback(self.expr(node.args[0]), walks=by_node.walks) - 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 @@ -304,9 +296,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 @@ -334,7 +326,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, @@ -356,7 +348,7 @@ def _partition_of(node: FunctionCallNode) -> program.Walk | None: return by_node.walks[0] -def _bound_expression(value: float | str) -> program.ExpressionNode: +def _bound_expression(value: float | str) -> program.Expression: if isinstance(value, str): return program.Parameter(value) return program.Constant(value) diff --git a/src/math_spec/program.py b/src/math_spec/program.py index 7567aa09..262cb065 100644 --- a/src/math_spec/program.py +++ b/src/math_spec/program.py @@ -22,7 +22,7 @@ from collections.abc import Mapping from dataclasses import dataclass, field, fields, replace from functools import cached_property -from typing import TYPE_CHECKING, Literal, NamedTuple, assert_never, get_args +from typing import TYPE_CHECKING, Literal, assert_never, get_args import math_spec.model as _model from math_spec._expression_parser import ComparisonOperator @@ -38,28 +38,26 @@ __all__ = [ 'QUADRATIC_POSITIONS', 'Add', - 'AndNode', - 'At', + 'And', 'AtLeastTwo', - 'BooleanLiteralNode', + 'BooleanLiteral', 'Cases', 'Check', - 'ConnectiveWhereNode', + 'Connective', 'Constant', 'ConstraintDeclaration', 'ConstraintSense', 'Contiguous', 'Curved', 'Derivation', - 'DimensionComparisonNode', + 'DimensionComparison', 'DimensionDeclaration', 'DimensionDtype', - 'DimensionPositionNode', + 'DimensionPosition', 'Divide', 'Dual', 'Expression', 'ExpressionDeclaration', - 'ExpressionNode', 'FanIn', 'FirstOf', 'Footprint', @@ -70,39 +68,40 @@ 'MaskOf', 'Multiply', 'Negate', - 'NotNode', + 'Not', 'ObjectiveDeclaration', 'ObjectiveSense', - 'OrNode', + 'Or', 'Parameter', - 'ParameterComparisonNode', + 'ParameterComparison', 'ParameterDeclaration', - 'ParameterDefinedNode', + 'ParameterDefined', 'ParameterDtype', 'PiecewiseDeclaration', 'Power', + 'Predicate', 'PredicateOperator', 'Program', + 'Pullback', 'QuadraticPosition', 'Reach', 'Region', - 'RelationComparisonNode', + 'RelationComparison', 'RelationDeclaration', - 'RelationDefinedNode', - 'RelationPairComparisonNode', + 'RelationDefined', + 'RelationPairComparison', 'Separability', 'SosDeclaration', 'Sum', 'Translate', - 'TypedPredicateNode', + 'TypedPredicate', 'Variable', 'VariableAbsence', 'VariableDeclaration', - 'VariableDefinedNode', + 'VariableDefined', 'VariableDomain', 'Walk', - 'WhereNode', - 'Window', + 'WindowSum', 'carries_variable', 'check_message', 'children', @@ -156,47 +155,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: @@ -210,30 +190,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``), @@ -241,28 +221,28 @@ class Power(Expression): coordinate like any other parameter arithmetic. """ - base: ExpressionNode - exponent: ExpressionNode + base: Expression + exponent: Expression @dataclass(frozen=True) -class Divide(Expression): +class Divide: """Quotient ``numerator / divisor``, the divisor variable-free wherever the math reads it (``math_spec.degree``).""" - numerator: ExpressionNode - divisor: ExpressionNode + numerator: Expression + divisor: Expression @dataclass(frozen=True) -class Sum(Expression): +class Sum: """Sum ``operand`` over the named dims, removing them from the result.""" - operand: ExpressionNode + operand: Expression over: tuple[str, ...] @dataclass(frozen=True) -class GroupSum(Expression): +class GroupSum: """Sum ``operand`` through relations, consuming the dims ``over`` and producing ``into``. ``walks`` says, per relation, which columns are consumed, which produced @@ -275,7 +255,7 @@ class GroupSum(Expression): produced column too where the operand already carries its dimension. """ - operand: ExpressionNode + operand: Expression walks: tuple[Walk, ...] @property @@ -292,7 +272,7 @@ def into(self) -> tuple[str, ...]: @dataclass(frozen=True) -class At(Expression): +class Pullback: """Read ``operand`` through relations — the adjoint of :class:`GroupSum`. Same tables, walked the other way: this consumes the dims in ``into`` and @@ -304,7 +284,7 @@ class At(Expression): :class:`GroupSum`, ``walks`` is the fact and the three are read off it. """ - operand: ExpressionNode + operand: Expression walks: tuple[Walk, ...] @property @@ -321,7 +301,7 @@ def into(self) -> tuple[str, ...]: @dataclass(frozen=True) -class Translate(Expression): +class Translate: """Re-index along one dimension: the result at *t* is ``operand`` at *t - offset*. ``wrap`` is ``edge='wrap'`` in the file: periodic, and stated on every @@ -330,9 +310,9 @@ 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 walked along ``dimension`` — its consumed + ``partition`` is a relation walked along ``along`` — its consumed column is a key over that dimension, its produced columns are the group — and the translation then happens inside each group: the neighbour is the one before in the same group, the edge is the group's, and a wrap closes @@ -340,8 +320,8 @@ class Translate(Expression): nothing. """ - operand: ExpressionNode - dimension: str + operand: Expression + along: str offset: int | str wrap: bool fill: float | None = None @@ -349,7 +329,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 @@ -369,8 +349,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: Walk | None = None @@ -385,11 +365,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 @@ -400,13 +380,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 @@ -418,14 +402,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 @@ -433,17 +417,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,) @@ -453,7 +437,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) @@ -467,8 +451,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; ``key`` is the roles a row is identified by, empty for a @@ -478,7 +463,6 @@ class RelationDeclaration(NamedTuple): that places them, and is what lets ``at`` read one value. """ - name: str columns: tuple[tuple[str, str], ...] key: tuple[str, ...] = () @@ -499,9 +483,11 @@ def dim(self, role: str) -> str: return dict(self.columns)[role] -class Walk(NamedTuple): +@dataclass(frozen=True) +class Walk: """One relation as an operator walks it — which columns are consumed, which produced, which joined on. + ``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 not walked (every role, for a bare relation): @@ -511,15 +497,12 @@ class Walk(NamedTuple): the group — every value role unless the call named some with ``within=``. """ + name: str relation: RelationDeclaration consumed: tuple[str, ...] produced: tuple[str, ...] joined: tuple[str, ...] - @property - def name(self) -> str: - return self.relation.name - @property def key(self) -> tuple[str, ...]: return self.relation.key @@ -556,9 +539,8 @@ def is_function_read(self) -> bool: @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 @@ -728,8 +710,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' @@ -744,9 +726,9 @@ class ConstraintDeclaration: """ dims: tuple[str, ...] - lhs: ExpressionNode + lhs: Expression sense: ConstraintSense - rhs: ExpressionNode + rhs: Expression where: Mask | None = None @@ -758,7 +740,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 @@ -776,7 +758,7 @@ class ObjectiveDeclaration: """Objective — scalar, every reduction in it one the file wrote.""" sense: ObjectiveSense - expression: ExpressionNode + expression: Expression @dataclass(frozen=True) @@ -785,13 +767,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 @@ -807,21 +789,13 @@ class Footprint: stands in; empty is affine throughout. domains: Every domain declared. sos_types: The order of each special-ordered set declared. - shapes: Every expression node kind that appears. + nodes: Every expression node kind that appears. """ quadratic: frozenset[QuadraticPosition] domains: frozenset[VariableDomain] sos_types: frozenset[Literal[1, 2]] - shapes: frozenset[type[ExpressionNode]] - - -def _declared[Declaration](items: Mapping[str, Declaration], name: str, kind: str) -> Declaration: - """The declaration called *name*, or a ``KeyError`` naming the near miss.""" - try: - return items[name] - except KeyError: - raise KeyError(f"unknown {kind} '{name}'. " + did_you_mean(name, list(items))) from None + nodes: frozenset[type[Expression]] @dataclass(frozen=True) @@ -940,6 +914,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`. @@ -949,7 +924,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.""" @@ -958,16 +933,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) @@ -980,23 +955,9 @@ def footprint(self) -> Footprint: ), domains=frozenset(v.domain for v in self.variables.values()), sos_types=frozenset(s.sos_type for s in self.sos.values()), - shapes=frozenset(type(node) for node in walk(*self.expressions)), + nodes=frozenset(type(node) for node in walk(*self.roots)), ) - def dimension(self, name: str) -> DimensionDeclaration: - return _declared(self.dimensions, name, 'dimension') - - @property - def 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. @@ -1027,7 +988,7 @@ def separability(self) -> Mapping[str, Separability]: # -------------------------------------------------------------------------- -def walk(*expressions: ExpressionNode) -> Iterator[ExpressionNode]: +def walk(*expressions: Expression) -> Iterator[Expression]: """Every node under *expressions*, each expression itself included, parents first. The traversal every *question* about a program is a filter of — which names @@ -1042,7 +1003,7 @@ def walk(*expressions: ExpressionNode) -> Iterator[ExpressionNode]: yield from walk(*children(expression)) -def is_quadratic(expression: ExpressionNode) -> bool: +def is_quadratic(expression: Expression) -> bool: """Whether *expression* contains a product of two variable-carrying operands. A structural question over the program, and unrelated consumers ask it — @@ -1061,22 +1022,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 @@ -1087,7 +1048,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))) @@ -1101,12 +1062,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 @@ -1119,7 +1080,7 @@ class ParameterDefinedNode: @dataclass(frozen=True) -class VariableDefinedNode: +class VariableDefined: """True at the coordinates where the named variable exists.""" name: str @@ -1127,7 +1088,7 @@ class VariableDefinedNode: @dataclass(frozen=True) -class ParameterComparisonNode: +class ParameterComparison: """Compare a parameter against a literal, element-wise.""" name: str @@ -1137,7 +1098,7 @@ class ParameterComparisonNode: @dataclass(frozen=True) -class DimensionComparisonNode: +class DimensionComparison: """Compare a dimension's own coordinates against a literal.""" name: str @@ -1146,7 +1107,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 @@ -1163,7 +1124,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 @@ -1178,7 +1139,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 @@ -1194,7 +1155,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 for a keyed @@ -1207,77 +1168,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. @@ -1285,9 +1246,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: @@ -1295,7 +1256,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 @@ -1308,19 +1269,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 @@ -1330,23 +1291,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 @@ -1355,12 +1310,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 @@ -1369,27 +1324,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 @@ -1406,14 +1361,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 @@ -1422,7 +1377,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) @@ -1443,12 +1398,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 e70d1aab..5edc4d80 100644 --- a/src/math_spec/resolution.py +++ b/src/math_spec/resolution.py @@ -61,24 +61,24 @@ unknown_operator_message, ) from math_spec.program import ( - AndNode, - BooleanLiteralNode, - DimensionComparisonNode, - DimensionPositionNode, + And, + BooleanLiteral, + DimensionComparison, + DimensionPosition, Mask, - NotNode, - OrNode, - ParameterComparisonNode, - ParameterDefinedNode, + Not, + Or, + ParameterComparison, + ParameterDefined, + Predicate, PredicateOperator, - RelationComparisonNode, + RelationComparison, RelationDeclaration, - RelationDefinedNode, - RelationPairComparisonNode, - TypedPredicateNode, - VariableDefinedNode, + RelationDefined, + RelationPairComparison, + TypedPredicate, + VariableDefined, Walk, - WhereNode, ) if TYPE_CHECKING: @@ -135,7 +135,7 @@ def of(cls, schema: Spec) -> Namespace: schema.variables, schema.parameters, schema.dimensions, - {n: RelationDeclaration(n, lk.pairs, lk.keys) for n, lk in schema.relations.items()}, + {n: RelationDeclaration(lk.pairs, lk.keys) 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()}, @@ -272,7 +272,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. @@ -291,9 +291,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) @@ -322,22 +322,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( @@ -346,7 +346,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: @@ -708,7 +708,7 @@ def _walk( ) return None joined = tuple(r for r in (shape.key or shape.roles) if r not in from_roles and r not in into_roles) - walk = Walk(shape, from_roles, into_roles, joined) + walk = Walk(name, shape, from_roles, into_roles, joined) if not forward and not walk.is_function_read: self.errors.append( f"{context}: {call}: at reads one value per coordinate, and '{name}' is not single-valued in " @@ -779,7 +779,7 @@ def _partition_walk( return None (walked,) = over_keys joined = tuple(r for r in shape.key if r != walked) - return Walk(shape, (walked,), shape.values if within_roles is None else within_roles, joined) + return Walk(name, shape, (walked,), shape.values if within_roles is None else within_roles, joined) def _default_role(self, name: str, call: str, kwarg: str, side: tuple[str, ...], what: str) -> str | None: """The one column *side* offers, or the refusal naming what the call has to choose from.""" @@ -823,9 +823,9 @@ 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) @@ -833,19 +833,19 @@ def where(self, node: WhereNode | UnresolvedWhereNode) -> WhereNode | Unresolved return self._position(node) if isinstance(node, UnresolvedComparisonNode): return self._comparison(node) - if isinstance(node, NotNode): - return NotNode(self._child(node.operand)) - if isinstance(node, AndNode): - return AndNode(self._child(node.left), self._child(node.right)) - if isinstance(node, OrNode): - return OrNode(self._child(node.left), self._child(node.right)) + if isinstance(node, Not): + return Not(self._child(node.operand)) + if isinstance(node, And): + return And(self._child(node.left), self._child(node.right)) + if isinstance(node, Or): + return Or(self._child(node.left), self._child(node.right)) assert_never(node) - def _child(self, node: WhereNode | UnresolvedWhereNode) -> WhereNode: + def _child(self, node: Predicate | UnresolvedWhereNode) -> Predicate: """A connective's child, typed as resolved: an unresolved one survives only with its refusal appended.""" - return cast('WhereNode', self.where(node)) + return cast('Predicate', self.where(node)) - def _where_name(self, node: UnresolvedNameNode) -> WhereNode | UnresolvedWhereNode: + def _where_name(self, node: UnresolvedNameNode) -> Predicate | UnresolvedWhereNode: """A bare name: a parameter's or relation's definedness, or a variable's existence.""" ns, context = self.ns, self.context kind = ns.kind(node.name) @@ -854,7 +854,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 " @@ -871,7 +871,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( @@ -880,10 +880,10 @@ def _where_name(self, node: UnresolvedNameNode) -> WhereNode | UnresolvedWhereNo f'exists. Test a parameter, or another variable declared before it.' ) else: - return VariableDefinedNode(node.name, ns.leaf_dims[node.name]) + return VariableDefined(node.name, ns.leaf_dims[node.name]) return node - def _position(self, node: UnresolvedPositionNode) -> DimensionPositionNode | UnresolvedPositionNode: + def _position(self, node: UnresolvedPositionNode) -> DimensionPosition | UnresolvedPositionNode: """``position(dim[, by=relation[, within=columns]]) i``: the name a dimension, ``by=`` a relation keyed over it.""" ns, context = self.ns, self.context if node.dimension not in ns.dimensions: @@ -894,7 +894,7 @@ def _position(self, node: UnresolvedPositionNode) -> DimensionPositionNode | Unr ) return node if node.by is None: - return DimensionPositionNode(node.dimension, node.op, node.position) + return DimensionPosition(node.dimension, node.op, node.position) call = f'position({node.dimension}, by={node.by})' if ns.kind(node.by) != 'relation': self.errors.append( @@ -906,9 +906,9 @@ def _position(self, node: UnresolvedPositionNode) -> DimensionPositionNode | Unr walk = self._partition_walk(node.by, 'position', node.dimension, node.into) if walk is None: return node - return DimensionPositionNode(node.dimension, node.op, node.position, walk) + return DimensionPosition(node.dimension, node.op, node.position, walk) - def _comparison(self, node: UnresolvedComparisonNode) -> WhereNode | UnresolvedWhereNode: + def _comparison(self, node: UnresolvedComparisonNode) -> Predicate | UnresolvedWhereNode: """``name literal``, or the one structural form ``relation relation``.""" ns, context = self.ns, self.context value = node.value @@ -925,7 +925,7 @@ def _comparison(self, node: UnresolvedComparisonNode) -> WhereNode | UnresolvedW self.errors.append(refusal) return node dims = tuple(ns.shape_of(left_name).dim(k) for k in ns.shape_of(left_name).key) - return RelationPairComparisonNode(left_name, left, right_name, right, node.op, dims) + return RelationPairComparison(left_name, left, right_name, right, node.op, dims) self.errors.append(_declared_rhs_error(context, node, value, rhs_kind)) return node @@ -957,13 +957,13 @@ def _comparison(self, node: UnresolvedComparisonNode) -> WhereNode | UnresolvedW match kind: case 'parameter': assert not isinstance(value, datetime.date) - return ParameterComparisonNode(left_name, node.op, value, ns.leaf_dims[left_name]) + return ParameterComparison(left_name, node.op, value, ns.leaf_dims[left_name]) case 'dimension': - return DimensionComparisonNode(left_name, node.op, value) + return DimensionComparison(left_name, node.op, value) case 'relation': assert column is not None shape = ns.shape_of(left_name) - return RelationComparisonNode(left_name, column, node.op, value, tuple(shape.dim(k) for k in shape.key)) + return RelationComparison(left_name, column, node.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 018a2ddd..cf0c70a0 100644 --- a/src/math_spec/separability.py +++ b/src/math_spec/separability.py @@ -9,26 +9,26 @@ from typing import TYPE_CHECKING, Literal from math_spec.program import ( - At, Cases, - DimensionPositionNode, + DimensionPosition, GroupSum, Mask, + Pullback, Reach, Separability, Sum, Translate, - Window, + WindowSum, walk, ) if TYPE_CHECKING: from collections.abc import Iterator - from math_spec.program import ExpressionNode, Program + from math_spec.program import Expression, Program -def _built_blocks(program: Program) -> Iterator[tuple[str, tuple[ExpressionNode, ...], Mask | None, bool]]: +def _built_blocks(program: Program) -> Iterator[tuple[str, tuple[Expression, ...], Mask | None, bool]]: """Every block that builds rows, labelled as the lowering's own messages label it. A named expression is not one: it is inlined where it is referenced, so @@ -91,12 +91,12 @@ def waits_on(dimension: str, label: str, name: str, kind: Literal['offset', 'par label, f'groups {dimension} into {", ".join(node.into)} — window that dimension instead, or cut only at the group edges', ) - elif isinstance(node, At): + elif isinstance(node, Pullback): for dimension in node.into: for relation in node.coordinate: waits_on(dimension, label, relation, 'coordinate') - elif isinstance(node, (Translate, Window)): - dimension = node.dimension + elif isinstance(node, (Translate, WindowSum)): + dimension = node.along if node.wrap: report( 'coupled', @@ -108,7 +108,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') @@ -116,7 +116,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 bb883a6e..3430fc98 100644 --- a/src/math_spec/typesetting/walk.py +++ b/src/math_spec/typesetting/walk.py @@ -35,21 +35,21 @@ ) from math_spec.dimensions import dims_of from math_spec.program import ( - AndNode, - BooleanLiteralNode, - DimensionComparisonNode, - DimensionPositionNode, + And, + BooleanLiteral, + DimensionComparison, + DimensionPosition, 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 @@ -295,7 +295,7 @@ def _value_read(self, name: str, column: str, ctx: _Context) -> str: keyed = self.format.joined([ctx.subscript(dict(lk.pairs)[k]) for k in lk.keys], '') return self.format.apply(self._column(name, column, len(lk.values) == 1), keyed) - def _position_group(self, node: DimensionPositionNode, ctx: _Context) -> str: + def _position_group(self, node: DimensionPosition, ctx: _Context) -> str: """The group a grouped position counts within: the relation's group columns read at the row's key.""" assert node.partition is not None walk = node.partition @@ -553,33 +553,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 ( @@ -587,22 +587,22 @@ 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 else self._position_group(node, ctx) place = self._position(ctx.subscript(node.name), grouping) ordinal = self._ordinal(node.name, node.position, grouping) return f'{place} {self._op(_PREDICATES[node.op])} {ordinal}', comparison - if isinstance(node, 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): lk = self.schema.relations[node.name] if lk.keys: keyed = self.format.joined([ctx.subscript(dict(lk.pairs)[k]) for k in lk.keys], '') @@ -611,18 +611,18 @@ def _where(self, node: WhereNode, ctx: _Context) -> tuple[str, int]: row = self.format.parenthesise(self.format.joined([ctx.subscript(d) for d in lk.dims], '')) return f'{row} {self._op("in")} {self.format.upright(node.name)}', 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 e4017957..839f3b88 100644 --- a/src/math_spec/validation.py +++ b/src/math_spec/validation.py @@ -37,7 +37,7 @@ from math_spec.expansion import expand, parse_and_expand, parse_template from math_spec.model import Spec from math_spec.operators import BUILTINS, unknown_operator_message -from math_spec.program import BooleanLiteralNode +from math_spec.program import BooleanLiteral from math_spec.resolution import ( Namespace, Resolved, @@ -52,7 +52,7 @@ from pathlib import Path from math_spec.model import ExpressionBlock - from math_spec.program import WhereNode + from math_spec.program import Predicate def to_spec(model: str | Path | dict[str, Any] | Spec) -> Spec: @@ -185,11 +185,11 @@ def _named( found = len(errors) arms: list[CaseArm] = [] - masks: dict[str, WhereNode] = {} + masks: dict[str, Predicate] = {} for case_name, case in block.cases.items(): arm_context = case_context(name, case_name) when = resolve_where_text(case.when, ns, arm_context, errors) - if isinstance(when, BooleanLiteralNode): + if isinstance(when, BooleanLiteral): errors.append(_constant_arm(arm_context, value=when.value)) elif when is not None: masks[case_name] = when diff --git a/tests/test_dimensions.py b/tests/test_dimensions.py index d1ddc19d..de878a9c 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 @@ -499,6 +499,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 3b62be64..2a18d1aa 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 84ddb1f8..dfb0e2d7 100644 --- a/tests/test_lowering.py +++ b/tests/test_lowering.py @@ -24,35 +24,35 @@ from math_spec.program import ( QUADRATIC_POSITIONS, Add, - AndNode, - At, - BooleanLiteralNode, + And, + BooleanLiteral, Cases, Constant, - DimensionComparisonNode, + DimensionComparison, DimensionDeclaration, Divide, Dual, - ExpressionNode, + Expression, Footprint, GroupSum, Mask, Multiply, Negate, - NotNode, - OrNode, + Not, + Or, Parameter, - ParameterComparisonNode, - ParameterDefinedNode, + ParameterComparison, + ParameterDefined, Power, Program, + Pullback, Region, RelationDeclaration, Sum, Translate, Variable, Walk, - Window, + WindowSum, children, divisor_parameters, fan_in, @@ -70,7 +70,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 @@ -83,11 +83,11 @@ } #: `lk` and `lk2` as `sum` walks them: key consumed, value produced, nothing joined. -LK = RelationDeclaration('lk', (('g', 'g'), ('h', 'h')), ('g',)) -LK2 = RelationDeclaration('lk2', (('g', 'g'), ('z', 'z')), ('g',)) -LK_WALK = Walk(LK, ('g',), ('h',), ()) -LK2_WALK = Walk(LK2, ('g',), ('z',), ()) -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_WALK = Walk('lk', LK, ('g',), ('h',), ()) +LK2_WALK = Walk('lk2', LK2, ('g',), ('z',), ()) +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 @@ -150,7 +150,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' @@ -181,33 +181,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', ), ], @@ -215,7 +215,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, ( @@ -229,7 +229,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): @@ -261,9 +261,8 @@ def test_a_lowered_mask_cannot_be_rewritten_in_place(dispatch_program): def test_a_lowered_where_is_a_mask_that_answers_from_its_root(dispatch_program): """The `where` a lowering carries is a `Mask`, and its questions are its root's. - A consumer asks the mask — `where.names_read`, `where.conjuncts` — the way it - asks a dimension `dimension.targets`, rather than reaching for a free function - with the raw node. + A consumer asks the mask — `where.names_read`, `where.conjuncts` — rather + than reaching for a free function with the raw node. """ (v,) = dispatch_program.variables.values() @@ -305,17 +304,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): @@ -333,32 +332,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' @@ -370,9 +369,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' @@ -397,7 +396,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): @@ -427,7 +426,7 @@ def test_a_power_lowers_to_a_node_of_its_own(dispatch_schema): ), pytest.param( 'at(r, by=lk)', - At(Variable('r'), walks=(Walk(LK, ('h',), ('g',), ()),)), + Pullback(Variable('r'), walks=(Walk('lk', LK, ('h',), ('g',), ()),)), id='a-pullback-walks-the-same-table-back', ), pytest.param( @@ -453,28 +452,28 @@ def test_a_power_lowers_to_a_node_of_its_own(dispatch_schema): offset=1, wrap=False, fill=0.0, - partition=Walk(LK, ('g',), ('h',), ()), + partition=Walk('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)', - Window( + WindowSum( Variable('p'), 'g', width=2, wrap=False, - partition=Walk(LK, ('g',), ('h',), ()), + partition=Walk('lk', LK, ('g',), ('h',), ()), ), id='a-window-stops-at-the-edges-of-the-relation-it-names', ), @@ -512,36 +511,34 @@ def test_a_relation_lowers_with_the_walk_each_call_takes(): ) 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'), walks=(Walk(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'), walks=(Walk('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.over, zonal.into, zonal.coordinate) == (('generator',), ('zone',), ('zone_of',)), ( 'the dims a consumer reads are read off the walk' ) assert program.constraints['history'].lhs == GroupSum( - Variable('p'), walks=(Walk(declared, ('snapshot',), ('zone',), ('generator',)),) + Variable('p'), walks=(Walk('zone_of', declared, ('snapshot',), ('zone',), ('generator',)),) ), 'the same table walked from its other key column' priced = program.constraints['priced'].rhs - assert priced == At(Parameter('price'), walks=(Walk(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'), walks=(Walk('zone_of', declared, ('zone',), ('generator',), ('snapshot',)),) + ), 'and its adjoint consumes the value column and produces the key column' + assert isinstance(priced, Pullback) assert (priced.over, priced.into) == (('generator',), ('zone',)), ( - 'an at produces the fine dims and consumes the coarse' + 'a pullback produces the fine dims and consumes the coarse' ) - 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' @@ -550,14 +547,14 @@ 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, walks=(Walk(component_of, ('component',), ('flow',), ()),)) + component_of = RelationDeclaration((('flow', 'flow'), ('component', 'component')), ('flow',)) + pulled = Pullback(quotient, walks=(Walk('component_of', component_of, ('component',), ('flow',), ()),)) assert divisor_parameters(pulled) == frozenset({'rate'}), 'the walk descends through `At`' assert divisor_parameters(Sum(pulled, ('flow',))) == frozenset({'rate'}), 'and through a `Sum` over it' @@ -577,12 +574,12 @@ def test_a_quotient_is_found_whole_so_its_two_halves_stay_paired(): left = Divide(Variable('x'), Parameter('rate')) right = Divide(Variable('y'), Parameter('loss')) - found = quotients(Sum(left + right, ('flow',))) + found = quotients(Sum(Add(left, right), ('flow',))) assert [(variables_of(q.numerator), q.divisor) for q in found] == [ (frozenset({'x'}), Parameter('rate')), (frozenset({'y'}), Parameter('loss')), ], 'each quotient keeps its own numerator, in the order the expression writes them' - assert divisor_parameters(Sum(left + right, ('flow',))) == frozenset({'rate', 'loss'}), ( + assert divisor_parameters(Sum(Add(left, right), ('flow',))) == frozenset({'rate', 'loss'}), ( 'the flat answer is still the union of the same walk' ) @@ -597,11 +594,11 @@ def test_a_quotient_is_found_whole_so_its_two_halves_stay_paired(): 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'), walks=(Walk(AT_BUS, ('g',), ('bus',), ()),)): 'many-to-one', - At(Variable('p'), walks=(Walk(AT_BUS, ('bus',), ('g',), ()),)): 'one-to-one', + GroupSum(Variable('p'), walks=(Walk('at_bus', AT_BUS, ('g',), ('bus',), ()),)): 'many-to-one', + Pullback(Variable('p'), walks=(Walk('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', } @@ -609,8 +606,8 @@ def test_a_quotient_is_found_whole_so_its_two_halves_stay_paired(): def test_every_expression_node_is_classified_by_fan_in(): """`fan_in` was a ClassVar on five nodes, so `Add(...).fan_in` was an AttributeError.""" covered = {type(node) for node in FAN_IN} - assert covered == set(get_args(ExpressionNode)), ( - 'every node in the ExpressionNode union is classified, and nothing retired lingers' + assert covered == set(get_args(Expression)), ( + 'every node in the Expression union is classified, and nothing retired lingers' ) @@ -619,38 +616,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': {'columns': ['g', 'season'], 'key': 'g'}, + 'at_bus': {'columns': ['g', 'bus'], 'key': 'g'}, + }, + ) ) - 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(): @@ -659,15 +643,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, @@ -676,16 +660,16 @@ def test_expressions_are_the_ones_a_row_is_built_from(): ) ) - assert list(program.named_expressions) == ['spend'], 'the declared ones keep their own name' - assert program.expressions == ( + assert list(program.expressions) == ['spend'], 'the declared ones keep their own name' + assert program.roots == ( program.objective.expression, program.constraints['c'].lhs, program.constraints['c'].rhs, ), 'the objective first, then both sides of each constraint, in declaration order' - assert program.named_expressions['spend'].expression not in program.expressions, ( + assert program.expressions['spend'].expression not in program.roots, ( 'a named expression builds no row, so it is not one of the expressions a row is built from' ) - assert len(program.expressions) == 3, 'and nothing else is counted' + assert len(program.roots) == 3, 'and nothing else is counted' def _footprint_of(constraint: str, objective: str) -> Footprint: @@ -724,7 +708,7 @@ def test_a_construct_the_file_does_not_use_is_an_empty_set_rather_than_none(): assert footprint.sos_types == frozenset(), 'a file declaring no sos' assert footprint.quadratic == frozenset(), 'a file with no quadratic anywhere' assert footprint.domains == {'continuous'}, 'never empty — a program has variables' - assert {type(f) for f in (footprint.sos_types, footprint.quadratic, footprint.shapes)} == {frozenset}, ( + assert {type(f) for f in (footprint.sos_types, footprint.quadratic, footprint.nodes)} == {frozenset}, ( 'every field is a set, so one rule reads all of them' ) assert footprint.quadratic <= QUADRATIC_POSITIONS, 'and the vocabulary a consumer pins its table against' @@ -742,8 +726,8 @@ def test_a_named_expression_is_not_in_the_footprint(): """It builds no row, so counting it would answer wrongly about what is solved.""" program = to_program(override(TINY, expressions={'spend': 'sum(p * cost, over=g)'})) - assert Parameter not in program.footprint.shapes, "the named expression's parameter reaches no row" - assert Parameter in {type(n) for n in walk(program.named_expressions['spend'].expression)}, ( + assert Parameter not in program.footprint.nodes, "the named expression's parameter reaches no row" + assert Parameter in {type(n) for n in walk(program.expressions['spend'].expression)}, ( 'though it is in the expression' ) @@ -756,8 +740,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 = { @@ -804,10 +788,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' ) @@ -844,10 +826,10 @@ def test_the_lowered_regions_are_still_proved_apart(): def test_a_cased_expression_is_readable_by_the_name_the_file_wrote(): - """`Program.expressions` carries it, so a consumer reads it back whole.""" + """`Program.expressions` carries it under its name, so a consumer reads it back whole.""" program = to_program(CASED) - assert isinstance(program.named_expressions['previous'].expression, Cases), ( + assert isinstance(program.expressions['previous'].expression, Cases), ( 'a cased expression reaches the program as the node, not as its fallback arm alone' ) @@ -881,7 +863,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(): @@ -893,7 +875,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' ) @@ -909,7 +891,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' ) @@ -917,7 +899,7 @@ def test_a_macro_formal_named_like_an_entry_keeps_the_entry_out_of_the_math(): def test_an_entry_that_reads_a_dual_is_a_reported_quantity(): """A dual is read after the solve, so an entry calling one is never in the math: it lowers to a Dual leaf and stays reported.""" program = to_program(override(TINY, expressions={'shadow_price': 'dual(c)'})) - declaration = program.named_expressions['shadow_price'] + declaration = program.expressions['shadow_price'] assert declaration.in_math is False, 'the entry reading a dual is reported, never in the math' assert isinstance(declaration.expression, Dual), 'and it lowers to a Dual leaf' diff --git a/tests/test_parser.py b/tests/test_parser.py index 3748a75f..74160a70 100644 --- a/tests/test_parser.py +++ b/tests/test_parser.py @@ -32,7 +32,7 @@ parse_where, ) from math_spec.errors import SchemaError -from math_spec.program import AndNode, BooleanLiteralNode, NotNode, OrNode, _conjuncts +from math_spec.program import And, BooleanLiteral, Not, Or, _conjuncts def test_the_grammar_builds_the_program_s_own_node_classes(): @@ -40,13 +40,13 @@ def test_the_grammar_builds_the_program_s_own_node_classes(): The parser constructs the resolved vocabulary's connectives directly, so a consumer's `isinstance` against the program's classes holds on any tree — - two homes for `AndNode` would make it hold on neither. + two homes for `And` would make it hold on neither. """ tree = parse_where('a AND NOT b OR True') - assert type(tree) is program_module.OrNode - assert type(tree.left) is program_module.AndNode - assert type(tree.left.right) is program_module.NotNode - assert type(tree.right) is program_module.BooleanLiteralNode + assert type(tree) is program_module.Or + assert type(tree.left) is program_module.And + assert type(tree.left.right) is program_module.Not + assert type(tree.right) is program_module.BooleanLiteral @pytest.mark.parametrize( @@ -240,12 +240,12 @@ def test_a_name_may_begin_with_inf(name): @pytest.mark.parametrize( ('text', 'node_type', 'attrs'), [ - pytest.param('True', BooleanLiteralNode, {'value': True}, id='a-literal'), + pytest.param('True', BooleanLiteral, {'value': True}, id='a-literal'), pytest.param('p_max', UnresolvedNameNode, {'name': 'p_max'}, id='a-bare-name'), pytest.param('p_max > 0', UnresolvedComparisonNode, {'op': '>', 'value': 0}, id='a-comparison'), - pytest.param('a AND b', AndNode, {}, id='and'), - pytest.param('a OR b', OrNode, {}, id='or'), - pytest.param('NOT a', NotNode, {}, id='not'), + pytest.param('a AND b', And, {}, id='and'), + pytest.param('a OR b', Or, {}, id='or'), + pytest.param('NOT a', Not, {}, id='not'), ], ) def test_a_where_string_parses_to_its_node(text, node_type, attrs): @@ -258,8 +258,8 @@ def test_a_where_string_parses_to_its_node(text, node_type, attrs): def test_and_binds_tighter_than_or(): - assert parse_where('a OR b AND c') == OrNode( - UnresolvedNameNode('a'), AndNode(UnresolvedNameNode('b'), UnresolvedNameNode('c')) + assert parse_where('a OR b AND c') == Or( + UnresolvedNameNode('a'), And(UnresolvedNameNode('b'), UnresolvedNameNode('c')) ) @@ -273,7 +273,7 @@ def test_and_binds_tighter_than_or(): ids=['single', 'pair', 'chain'], ) def test_conjuncts_flattens_the_and_spine(text, expected): - """A chain the grammar left-folds into nested `AndNode`s comes back flat (#312). + """A chain the grammar left-folds into nested `And`s comes back flat (#312). `_conjuncts` is the one home of the flatten rule; `Mask.conjuncts` is the door a consumer asks it through.""" @@ -336,7 +336,7 @@ def test_position_converts_a_dimension_to_where_a_row_sits(text, op, position, b def test_a_position_is_not_confused_with_a_name(): """`position` leads the alternation, so it is not read as a bare name.""" - assert isinstance(parse_where('position(t) == 0 AND p_max > 0'), AndNode) + assert isinstance(parse_where('position(t) == 0 AND p_max > 0'), And) @pytest.mark.parametrize( diff --git a/tests/test_piecewise.py b/tests/test_piecewise.py index 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 cca1d1c8..7ba71bad 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' @@ -461,7 +461,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 e93b0712..0669ec07 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 @@ -146,10 +146,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 walk 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', 'Walk', 'RelationDeclaration'} def test_the_golden_model_carries_every_node_kind_the_walk_renders(): @@ -161,7 +162,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, '