Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 5 additions & 5 deletions docs/contributing.md
Original file line number Diff line number Diff line change
Expand Up @@ -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:

Expand Down
4 changes: 2 additions & 2 deletions docs/reference/language/reading.md
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand 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
Expand Down
2 changes: 1 addition & 1 deletion docs/reference/language/reported.md
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
6 changes: 3 additions & 3 deletions src/math_spec/_expression_parser.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down Expand Up @@ -191,7 +191,7 @@ class CaseArm:
"""

label: str
when: WhereNode | None
when: Predicate | None
value: ArithmeticNode


Expand Down Expand Up @@ -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

Expand Down
24 changes: 12 additions & 12 deletions src/math_spec/_where_parser.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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]))
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand All @@ -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),
)
6 changes: 3 additions & 3 deletions src/math_spec/advice.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
12 changes: 6 additions & 6 deletions src/math_spec/boundedness.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
)
Expand Down Expand Up @@ -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
Expand All @@ -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
Expand Down Expand Up @@ -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
Expand Down
24 changes: 12 additions & 12 deletions src/math_spec/dimensions.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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)
Expand Down
Loading
Loading