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
4 changes: 2 additions & 2 deletions docs/about/relations-as-linear-maps.md
Original file line number Diff line number Diff line change
Expand Up @@ -131,8 +131,8 @@ $`x`$ over generators and $`y`$ over buses,
\langle M x, y \rangle = \sum_{b} y_b \sum_{g} \mathbf{1}_R(g, b)\, x_g = \sum_{g} x_g \sum_{b} \mathbf{1}_R(g, b)\, y_b = \langle x, M^{\mathsf{T}} y \rangle,
```

which is why the program's `Lookup` is the same `Join` as its `GroupSum`, read
without the group-by. A bare relation has the same matrix without the
which is why the program lowers `at` to a `Join` node and `sum(by=)` to the
same `Join` under a `Sum`. A bare relation has the same matrix without the
functional claim. A column of $`M`$ may hold several ones, so the sum fans out
and no group is one row. That is why `at` through a bare relation is refused.

Expand Down
14 changes: 7 additions & 7 deletions docs/contributing.md
Original file line number Diff line number Diff line change
Expand Up @@ -107,13 +107,13 @@ suffix says which layer:
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 join and group-by |
| `at(by=)` | `Lookup` | a join with no group-by |
| `shift(along=)` | `Translate` | a re-index along one dimension |
| `sum_back(along=)` | `WindowSum` | a sum over a trailing window |
| File verb | Node | What the node names |
| ------------------ | ------------------- | ------------------------------------------ |
| `sum(over=)` | `Sum` | dims removed from the result |
| `sum(by=)` | `Sum` over a `Join` | a join, and the sum over the dims it drops |
| `at(by=)` | `Join` | a join with no sum over it |
| `shift(along=)` | `Translate` | a re-index along one dimension |
| `sum_back(along=)` | `WindowSum` | a sum over a trailing window |

Nothing is abbreviated.

Expand Down
8 changes: 4 additions & 4 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 Join, Partition, Predicate
from math_spec.program import JoinColumns, 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 @@ -133,12 +133,12 @@ def __str__(self) -> str:

@dataclass(frozen=True)
class JoinNode:
"""A resolved ``by=`` on ``sum`` or ``at``: the relation, as the :class:`Join` the call names."""
"""A resolved ``by=`` on ``sum`` or ``at``: the relation, as the :class:`JoinColumns` the call names."""

join: Join
columns: JoinColumns

def __str__(self) -> str:
return self.join.name
return self.columns.name


@dataclass(frozen=True)
Expand Down
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 GroupSum, Lookup, walk
from math_spec.program import Join, walk

if TYPE_CHECKING:
from collections.abc import Mapping
Expand Down Expand Up @@ -74,6 +74,6 @@ def _grouped_axes(program: Program) -> set[str]:
"""
axes: set[str] = set()
for node in walk(*program.roots):
if isinstance(node, GroupSum | Lookup):
axes.update(node.join.added_dims)
if isinstance(node, Join):
axes.update(node.columns.added_dims)
return axes
5 changes: 2 additions & 3 deletions src/math_spec/boundedness.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,8 +24,7 @@
Divide,
Dual,
Expression,
GroupSum,
Lookup,
Join,
Multiply,
Negate,
Parameter,
Expand Down Expand Up @@ -162,7 +161,7 @@ def _record_signs(node: Expression, sign: Sign, signs: dict[str, Sign]) -> None:
_record_signs(node.base, None, signs)
_record_signs(node.exponent, None, signs)
return
if isinstance(node, Sum | GroupSum | Lookup | Translate | WindowSum | Cases):
if isinstance(node, Sum | Join | Translate | WindowSum | Cases):
for child in children(node):
_record_signs(child, sign, signs)
return
Expand Down
10 changes: 5 additions & 5 deletions src/math_spec/dimensions.py
Original file line number Diff line number Diff line change
Expand Up @@ -46,7 +46,7 @@
DimensionComparison,
DimensionPosition,
ExpressionComparison,
Join,
JoinColumns,
Mask,
ParameterComparison,
ParameterDefined,
Expand Down Expand Up @@ -161,7 +161,7 @@ def _sum_dims(node: FunctionCallNode, inner: frozenset[str], schema: Spec, conte
return inner - {summed.name}

assert isinstance(by, JoinNode), 'resolution reads sum(by=) as a join'
join = by.join
join = by.columns
if missing := sorted(set(join.dropped_dims) - inner):
raise DimensionError(
_not_carried(
Expand All @@ -178,7 +178,7 @@ def _at_dims(node: FunctionCallNode, inner: frozenset[str], schema: Spec, contex
"""``at`` is the join of ``sum(by=)`` with no group-by: it joins on the dims a sum groups by and groups by the ones a sum joins on."""
by = node.kwargs['by']
assert isinstance(by, JoinNode), 'resolution reads at(by=) as a join'
join = by.join
join = by.columns
if absent := sorted(set(join.dropped_dims) - inner):
raise DimensionError(
f'{context}: at(by={join.name}) joins on '
Expand Down Expand Up @@ -212,7 +212,7 @@ def _translation_dims(node: FunctionCallNode, inner: frozenset[str], schema: Spe
return inner


def _join_dims(call: str, join: Join, inner: frozenset[str], context: str) -> frozenset[str]:
def _join_dims(call: str, join: JoinColumns, inner: frozenset[str], context: str) -> frozenset[str]:
"""The dims after *join*: the operand's, less the dims joined on, plus the dims grouped by.

A column grouped by and not joined on brings its dim, so the operand does
Expand All @@ -234,7 +234,7 @@ def _join_dims(call: str, join: Join, inner: frozenset[str], context: str) -> fr
return (inner - set(join.joined_dims)) | set(join.grouped_dims)


def _check_joined(call: str, use: Join | Partition, inner: frozenset[str], context: str) -> None:
def _check_joined(call: str, use: JoinColumns | Partition, inner: frozenset[str], context: str) -> None:
"""The columns a call joins on are matched at their dimensions, so the operand carries every one, each once.

Two joined columns over one dimension would match the operand's one
Expand Down
13 changes: 7 additions & 6 deletions src/math_spec/lowering.py
Original file line number Diff line number Diff line change
Expand Up @@ -304,9 +304,9 @@ def _predicate(self, node: program.Predicate) -> program.Predicate:
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
*into* another are different relational shapes, so ``by=`` decides which
before anything else is read.
One node either way: a ``sum(by=)`` is a :class:`~math_spec.program.Sum`
over a :class:`~math_spec.program.Join`, its ``over`` the dims the join
drops, so ``by=`` decides only what the sum stands over.
"""
by_node = node.kwargs.get('by')
operand = self.expr(node.args[0])
Expand All @@ -317,13 +317,14 @@ def sum(self, node: FunctionCallNode) -> program.Expression:
assert isinstance(summed, DimensionNode), 'resolution refuses a over= that is not a dimension'
return program.Sum(operand, (summed.name,))
assert isinstance(by_node, JoinNode), 'resolution reads sum(by=) as a join'
return program.GroupSum(operand, join=by_node.join)
columns = by_node.columns
return program.Sum(program.Join(operand, columns), columns.dropped_dims)

def at(self, node: FunctionCallNode) -> program.Expression:
"""``at(x, by=relation)`` — the join of :meth:`sum`'s ``by=`` form with no group-by."""
"""``at(x, by=relation)`` — the :class:`~math_spec.program.Join` of :meth:`sum`'s ``by=`` form, with no sum over it."""
by_node = node.kwargs['by']
assert isinstance(by_node, JoinNode), 'resolution reads at(by=) as a join'
return program.Lookup(self.expr(node.args[0]), join=by_node.join)
return program.Join(self.expr(node.args[0]), by_node.columns)

def sum_back(self, node: FunctionCallNode) -> program.Expression:
"""``sum_back(x, along=d, window=w)`` — a trailing window along one dimension.
Expand Down
56 changes: 24 additions & 32 deletions src/math_spec/program.py
Original file line number Diff line number Diff line change
Expand Up @@ -64,11 +64,10 @@
'FanIn',
'FirstOf',
'Footprint',
'GroupSum',
'Holds',
'Join',
'JoinColumns',
'LastOf',
'Lookup',
'Mask',
'MaskOf',
'Multiply',
Expand Down Expand Up @@ -241,38 +240,32 @@ class Divide:

@dataclass(frozen=True)
class Sum:
"""Sum ``operand`` over the named dims, removing them from the result."""
"""Sum ``operand`` over the named dims, removing them from the result.

operand: Expression
over: tuple[str, ...]


@dataclass(frozen=True)
class GroupSum:
"""Sum ``operand`` through a relation: a join on ``join.joined``, then a group-by on ``join.grouped`` with a sum.

The operand carries every dim joined on. The result drops the dims joined
on and not grouped by, keeps the ones both joined on and grouped by, and
gains the ones grouped by and not joined on.
``sum(x, by=relation, over=a, into=b)`` lowers to a ``Sum`` over a
:class:`Join`, ``over`` naming the dims the join drops: the group-by is
this node, and the join is its operand.
"""

operand: Expression
join: Join
over: tuple[str, ...]


@dataclass(frozen=True)
class Lookup:
"""Read ``operand`` through a relation: the join of :class:`GroupSum` with no group-by.

The grouped columns hold the relation's whole key, which the loader
checks, so each row of the result meets one row of the relation and
reads one value. The join fans out where several key tuples share the
values joined on, at each coordinate of the columns both joined on and
grouped by.
class Join:
"""Join ``operand`` to a relation on the columns ``columns`` joins on, and carry the columns it groups by.

The operand carries every dim joined on. The result keeps every dim the
operand carries and gains the dims grouped by and not joined on; the dims
joined on and not grouped by leave only under a :class:`Sum` over them.
``at(x, by=relation, over=a, into=b)`` lowers to a bare ``Join``: the
grouped columns hold the relation's whole key, which the loader checks, so
each row of the result meets one row of the relation and reads one value.
The join fans out where several key tuples share the values joined on.
"""

operand: Expression
join: Join
columns: JoinColumns


@dataclass(frozen=True)
Expand Down Expand Up @@ -375,8 +368,7 @@ class Cases:
| Power
| Divide
| Sum
| GroupSum
| Lookup
| Join
| Translate
| WindowSum
| Cases
Expand All @@ -389,13 +381,13 @@ def fan_in(expression: Expression) -> FanIn:
For the absence rules, both classes other than ``'one-to-one'`` sum
several input slots into an output row.
"""
if isinstance(expression, (Sum, GroupSum)):
if isinstance(expression, Sum):
return 'many-to-one'
if isinstance(expression, WindowSum):
return 'one-to-many'
if isinstance(
expression,
(Constant, Parameter, Variable, Dual, Negate, Add, Multiply, Power, Divide, Lookup, Translate, Cases),
(Constant, Parameter, Variable, Dual, Negate, Add, Multiply, Power, Divide, Join, Translate, Cases),
):
return 'one-to-one'
assert_never(expression)
Expand All @@ -411,7 +403,7 @@ def children(expression: Expression) -> tuple[Expression, ...]:
return (expression.numerator, expression.divisor)
if isinstance(expression, Power):
return (expression.base, expression.exponent)
if isinstance(expression, (Sum, GroupSum, Lookup, Translate, WindowSum)):
if isinstance(expression, (Sum, Join, Translate, WindowSum)):
return (expression.operand,)
if isinstance(expression, Cases):
return tuple(region.value for region in expression.regions)
Expand Down Expand Up @@ -465,7 +457,7 @@ def dim(self, role: str) -> str:


@dataclass(frozen=True)
class Join:
class JoinColumns:
"""One relation as one call joins it: the columns joined on, and the columns grouped by.

The declaration fixes no direction; the call does, and this is the one it
Expand Down Expand Up @@ -1440,8 +1432,8 @@ def _names_under(*expressions: Expression) -> frozenset[str]:
for node in walk(*expressions):
if isinstance(node, Cases):
names.update(*(region.when.names_read for region in node.regions))
elif isinstance(node, (GroupSum, Lookup)):
names.add(node.join.name)
elif isinstance(node, Join):
names.add(node.columns.name)
elif isinstance(node, (Translate, WindowSum)):
if node.partition is not None:
names.add(node.partition.name)
Expand Down
12 changes: 6 additions & 6 deletions src/math_spec/resolution.py
Original file line number Diff line number Diff line change
Expand Up @@ -74,7 +74,7 @@
CountComparison,
DimensionComparison,
DimensionPosition,
Join,
JoinColumns,
Mask,
Not,
Or,
Expand Down Expand Up @@ -227,7 +227,7 @@ class Resolved:
objective: The objective's expression, ``None`` where the file
declares none.
relations: Each relation's columns and key, as declared — the one
copy, which every :class:`~math_spec.program.Join` and
copy, which every :class:`~math_spec.program.JoinColumns` and
:class:`~math_spec.program.Partition` in the trees holds.
assumptions: Each ``assumptions:`` entry's predicate and the mask it
is checked under.
Expand Down Expand Up @@ -594,8 +594,8 @@ def _relation_ref(
return value if partition is None else PartitionNode(partition)
if not ({'over', 'into'} <= set(named)):
return value # the call shape refused it already, with the wording that names the rewrite
join = self._join(name, operator, named['over'], named['into'])
return value if join is None else JoinNode(join)
columns = self._join(name, operator, named['over'], named['into'])
return value if columns is None else JoinNode(columns)

def _role_name(self, value: ArithmeticNode, operator: str, key: str) -> tuple[str, ...] | None:
"""``over=`` or ``into=`` as the column names it must be — one bare name, or a bracketed list of them."""
Expand All @@ -614,7 +614,7 @@ def _join(
operator: str,
from_roles: tuple[str, ...],
into_roles: tuple[str, ...],
) -> Join | None:
) -> JoinColumns | None:
"""How ``sum`` or ``at`` joins relation *name*, between the columns the call named.

Both ends arrive written: the call shape refuses a call that leaves
Expand Down Expand Up @@ -657,7 +657,7 @@ def _join(
)
return None
kept = tuple(r for r in shape.key if r not in from_roles and r not in into_roles)
join = Join(name, shape, (*from_roles, *kept), (*into_roles, *kept))
join = JoinColumns(name, shape, (*from_roles, *kept), (*into_roles, *kept))
one_row_per_group = set(shape.key) <= set(join.grouped)
if not forward and not one_row_per_group:
self.errors.append(
Expand Down
30 changes: 17 additions & 13 deletions src/math_spec/separability.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,8 +11,7 @@
from math_spec.program import (
Cases,
DimensionPosition,
GroupSum,
Lookup,
Join,
Mask,
Reach,
Separability,
Expand Down Expand Up @@ -94,9 +93,20 @@ def waits_on(dimension: str, label: str, name: str, kind: Literal['offset', 'par
if row is not None:
rows[label] = row
masks: list[Mask | None] = [mask]
grouped = {
id(node.operand) for node in walk(*nodes) if isinstance(node, Sum) and isinstance(node.operand, Join)
}
for node in walk(*nodes):
if isinstance(node, Cases):
masks.extend(region.when for region in node.regions)
elif isinstance(node, Sum) and isinstance(node.operand, Join):
for dimension in node.over:
report(
'coupled',
dimension,
label,
f'groups {dimension} into {", ".join(node.operand.columns.added_dims)} — window that dimension instead, or cut only at the group edges',
)
elif isinstance(node, Sum):
if reductions_couple:
for dimension in node.over:
Expand All @@ -106,17 +116,11 @@ def waits_on(dimension: str, label: str, name: str, kind: Literal['offset', 'par
label,
f'sums over {dimension} — a rolling sum_back(window=n) windows, a total over the horizon does not',
)
elif isinstance(node, GroupSum):
for dimension in node.join.dropped_dims:
report(
'coupled',
dimension,
label,
f'groups {dimension} into {", ".join(node.join.added_dims)} — window that dimension instead, or cut only at the group edges',
)
elif isinstance(node, Lookup):
for dimension in node.join.dropped_dims:
waits_on(dimension, label, node.join.name, 'coordinate')
elif isinstance(node, Join):
if id(node) in grouped:
continue
for dimension in node.columns.dropped_dims:
waits_on(dimension, label, node.columns.name, 'coordinate')
elif isinstance(node, (Translate, WindowSum)):
dimension = node.along
if node.wrap:
Expand Down
Loading
Loading