Skip to content
Merged
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
26 changes: 18 additions & 8 deletions docs/contributing.md
Original file line number Diff line number Diff line change
Expand Up @@ -98,14 +98,24 @@ stale anchor fails it. `pixi run docs-serve` builds the site and serves it at
The same construct passes through three layers, and each names it in full. The
suffix says which layer:

| Layer | Suffix | Example |
| ------------------------------- | -------------------- | ----------------------------------------- |
| YAML block (`math_spec.model`) | `Block` | `VariableBlock`, `PiecewiseBlock` |
| Core AST (`math_spec.*_parser`) | `Node` | `VariableNode`, `DimensionComparisonNode` |
| Program (`math_spec.program`) | none / `Declaration` | `Variable`, `VariableDeclaration` |

A node names the coordinate map rather than the operator: the translation node
is `Translate`, and the operator is `shift`. Nothing is abbreviated.
| Layer | Suffix | Example |
| ------------------------------- | -------------------- | ------------------------------------------ |
| YAML block (`math_spec.model`) | `Block` | `VariableBlock`, `PiecewiseBlock` |
| Core AST (`math_spec.*_parser`) | `Node` | `VariableNode`, `UnresolvedComparisonNode` |
| Program (`math_spec.program`) | none / `Declaration` | `Variable`, `VariableDeclaration` |

A node names the operation, not the verb a file writes. One verb can lower to
two nodes, so the file's spelling cannot decide the name.

| File verb | Node | What the node names |
| ------------------ | ----------- | ------------------------------ |
| `sum(over=)` | `Sum` | dims removed from the result |
| `sum(by=)` | `GroupSum` | a sum through a relation |
| `at(by=)` | `Pullback` | a read through a relation |
| `shift(along=)` | `Translate` | a re-index along one dimension |
| `sum_back(along=)` | `WindowSum` | a sum over a trailing window |

Nothing is abbreviated.

## Adding an operator

Expand Down
2 changes: 1 addition & 1 deletion docs/reference/reading.md
Original file line number Diff line number Diff line change
Expand Up @@ -110,7 +110,7 @@ footprint = program.footprint
sorted(footprint.quadratic) # []
sorted(footprint.domains) # ['continuous']
sorted(footprint.sos_types) # []
sorted(kind.__name__ for kind in footprint.shapes) # ['Constant', 'Multiply', 'Parameter', 'Sum', 'Variable']
sorted(kind.__name__ for kind in footprint.kinds) # ['Constant', 'Multiply', 'Parameter', 'Sum', 'Variable']
```

Every field is a set. An empty field means this model does not use the
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 @@ -23,7 +23,7 @@
if TYPE_CHECKING:
from collections.abc import Callable, Iterator, Mapping

from math_spec.program import Direction, Partition, WhereNode
from math_spec.program import Direction, Partition, Predicate

#: The relation a comparison may carry — the three an expression may be
#: written with, which is what a constraint's sense is read off.
Expand Down Expand Up @@ -222,7 +222,7 @@ class CaseArm:
"""

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


Expand Down Expand Up @@ -307,7 +307,7 @@ def __str__(self) -> str:


#: A whole spec-side expression tree — parse output and the resolved tree alike.
#: Named apart from :data:`math_spec.program.ExpressionNode`, the lowered
#: Named apart from :data:`math_spec.program.Expression`, the lowered
#: vocabulary a consumer reads.
ParsedNode = ArithmeticNode | ComparisonNode

Expand Down
38 changes: 19 additions & 19 deletions src/math_spec/_where_parser.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,13 +22,13 @@

from math_spec._expression_parser import ARITHMETIC, NAME, ArithmeticNode, children, parse_text
from math_spec.program import (
AndNode,
BooleanLiteralNode,
ConnectiveWhereNode,
NotNode,
OrNode,
And,
BooleanLiteral,
Connective,
Not,
Or,
Predicate,
PredicateOperator,
WhereNode,
where_children,
)

Expand Down Expand Up @@ -93,7 +93,7 @@ class UnresolvedComparisonNode:
#: Every node a parsed where string is built of: the connectives and literals,
#: the unresolved leaves, and the arithmetic and the two side nodes under a
#: comparison. What the depth measurement walks.
_ParsedWhere = WhereNode | UnresolvedWhereNode | ArithmeticNode | ColumnNode | QuotedNode
_ParsedWhere = Predicate | UnresolvedWhereNode | ArithmeticNode | ColumnNode | QuotedNode


# ---------------------------------------------------------------------------
Expand All @@ -110,8 +110,8 @@ def _build_where_grammar() -> pp.ParserElement:
"""
where_expr = pp.Forward()

true_lit = pp.CaselessKeyword('True').set_parse_action(lambda: BooleanLiteralNode(True))
false_lit = pp.CaselessKeyword('False').set_parse_action(lambda: BooleanLiteralNode(False))
true_lit = pp.CaselessKeyword('True').set_parse_action(lambda: BooleanLiteral(True))
false_lit = pp.CaselessKeyword('False').set_parse_action(lambda: BooleanLiteral(False))

name = pp.Regex(NAME)
# pyrefly: ignore[implicit-any-lambda]
Expand All @@ -133,28 +133,28 @@ def _build_where_grammar() -> pp.ParserElement:

NOT = pp.CaselessKeyword('NOT').suppress()
# pyrefly: ignore[implicit-any-lambda]
not_expr = (NOT + atom).set_parse_action(lambda t: NotNode(t[0])) | atom
not_expr = (NOT + atom).set_parse_action(lambda t: Not(t[0])) | atom

AND = pp.CaselessKeyword('AND').suppress()
and_expr = not_expr + pp.ZeroOrMore(AND + not_expr)
and_expr.set_parse_action(_folder(AndNode))
and_expr.set_parse_action(_folder(And))

OR = pp.CaselessKeyword('OR').suppress()
or_expr = and_expr + pp.ZeroOrMore(OR + and_expr)
or_expr.set_parse_action(_folder(OrNode))
or_expr.set_parse_action(_folder(Or))

where_expr <<= or_expr
return where_expr


def _folder(node_type: type[AndNode] | type[OrNode]) -> Callable[[pp.ParseResults], WhereNode | UnresolvedWhereNode]:
def _folder(node_type: type[And] | type[Or]) -> Callable[[pp.ParseResults], Predicate | UnresolvedWhereNode]:
"""A parse action left-folding a flat operator chain into *node_type*."""

def fold(tokens: pp.ParseResults) -> WhereNode | UnresolvedWhereNode:
items: list[WhereNode | UnresolvedWhereNode] = list(tokens)
def fold(tokens: pp.ParseResults) -> Predicate | UnresolvedWhereNode:
items: list[Predicate | UnresolvedWhereNode] = list(tokens)
result = items[0]
for item in items[1:]:
result = node_type(cast('WhereNode', result), cast('WhereNode', item))
result = node_type(cast('Predicate', result), cast('Predicate', item))
return result

return fold
Expand Down Expand Up @@ -196,13 +196,13 @@ def _nested(node: _ParsedWhere) -> tuple[_ParsedWhere, ...]:
return (node.left, node.right)
if isinstance(node, ArithmeticNode):
return children(node)
if isinstance(node, ConnectiveWhereNode):
if isinstance(node, Connective):
return where_children(node)
return ()


@lru_cache(maxsize=4096)
def parse_where(text: str) -> WhereNode | UnresolvedWhereNode:
def parse_where(text: str) -> Predicate | UnresolvedWhereNode:
"""Parse a where string into an AST, its leaves still unresolved.

The connectives and literals are the resolved vocabulary's own; the leaves
Expand All @@ -217,6 +217,6 @@ def parse_where(text: str) -> WhereNode | UnresolvedWhereNode:
as an expression is.
"""
return cast(
'WhereNode | UnresolvedWhereNode',
'Predicate | UnresolvedWhereNode',
parse_text(_WHERE_GRAMMAR, text, 'where string', _named_rewrite, _nested, _DEEP_REWRITE),
)
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 collections.abc import Mapping
Expand Down Expand Up @@ -73,7 +73,7 @@ def _produced_axes(program: Program) -> set[str]:
dimension: either way, the dims the direction produces.
"""
axes: set[str] = set()
for node in walk(*program.expressions):
if isinstance(node, GroupSum | At):
for node in walk(*program.roots):
if isinstance(node, GroupSum | Pullback):
axes.update(node.direction.produced_dims)
return axes
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 @@ -41,17 +41,17 @@
from math_spec.errors import DimensionError
from math_spec.operators import BUILTINS
from math_spec.program import (
DimensionComparisonNode,
DimensionPositionNode,
DimensionComparison,
DimensionPosition,
Direction,
Mask,
ParameterComparisonNode,
ParameterDefinedNode,
ParameterComparison,
ParameterDefined,
Partition,
RelationComparisonNode,
RelationDefinedNode,
RelationPairComparisonNode,
VariableDefinedNode,
RelationComparison,
RelationDefined,
RelationPairComparison,
VariableDefined,
)

if TYPE_CHECKING:
Expand Down Expand Up @@ -539,13 +539,13 @@ def _check_where_dims(
if not (outside := sorted(Mask(atom).dims - frame)):
continue
match atom:
case ParameterDefinedNode() | ParameterComparisonNode():
case ParameterDefined() | ParameterComparison():
noun = 'parameter'
case VariableDefinedNode():
case VariableDefined():
noun = 'variable'
case DimensionComparisonNode() | DimensionPositionNode():
case DimensionComparison() | DimensionPosition():
noun = 'dimension'
case RelationComparisonNode() | RelationPairComparisonNode() | RelationDefinedNode():
case RelationComparison() | RelationPairComparison() | RelationDefined():
noun = 'relation'
case _:
assert_never(atom)
Expand Down
Loading
Loading