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: 21 additions & 5 deletions src/math_spec/program.py
Original file line number Diff line number Diff line change
Expand Up @@ -267,10 +267,11 @@ class GroupSum(Expression):

``walks`` says, per relation, which columns are consumed, which produced
and which joined on, and is the one fact the node holds: ``coordinate``
names the relations, ``over`` is the dims every walk consumes and ``into``
the dims they produce, in walk order, so that several coordinates are
one grouping into a product of targets, consumed in a single join. The
result replaces every dim in ``over`` with every dim in ``into``. The
names the relations, ``over`` is the dims every walk consumes, ``into``
the dims they produce, in walk order, and ``joined`` the dims they join
on, so that several coordinates are one grouping into a product of
targets, consumed in a single join. The result replaces every dim in
``over`` with every dim in ``into`` and keeps every dim in ``joined``. The
join keys on the consumed columns and every joined column, and on a
produced column too where the operand already carries its dimension.
"""
Expand All @@ -290,6 +291,11 @@ def over(self) -> tuple[str, ...]:
def into(self) -> tuple[str, ...]:
return tuple(dim for walk in self.walks for dim in walk.produced_dims)

@property
def joined(self) -> tuple[str, ...]:
"""The dims the walks join on, each once — the key columns neither consumed nor produced, which the operand carries."""
return _joined_dims(self.walks)


@dataclass(frozen=True)
class At(Expression):
Expand All @@ -301,7 +307,7 @@ class At(Expression):
(``Walk.is_function_read``). The join fans out, many ``over`` tuples
sharing one ``into`` tuple — at each coordinate of the joined columns,
which the operand carries and the result keeps. As on
:class:`GroupSum`, ``walks`` is the fact and the three are read off it.
:class:`GroupSum`, ``walks`` is the fact and the four are read off it.
"""

operand: ExpressionNode
Expand All @@ -319,6 +325,16 @@ def over(self) -> tuple[str, ...]:
def into(self) -> tuple[str, ...]:
return tuple(dim for walk in self.walks for dim in walk.consumed_dims)

@property
def joined(self) -> tuple[str, ...]:
"""The dims the walks join on, each once — the key columns neither consumed nor produced, which the operand carries."""
return _joined_dims(self.walks)


def _joined_dims(walks: tuple[Walk, ...]) -> tuple[str, ...]:
"""The dims *walks* join on, each once, in walk order — the rule :attr:`GroupSum.joined` and :attr:`At.joined` share."""
return tuple(dict.fromkeys(dim for walk in walks for dim in walk.joined_dims))


@dataclass(frozen=True)
class Translate(Expression):
Expand Down
13 changes: 8 additions & 5 deletions tests/test_lowering.py
Original file line number Diff line number Diff line change
Expand Up @@ -521,9 +521,12 @@ def test_a_relation_lowers_with_the_walk_each_call_takes():
'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 (zonal.over, zonal.into, zonal.joined, zonal.coordinate) == (
('generator',),
('zone',),
('snapshot',),
('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',)),)
), 'the same table walked from its other key column'
Expand All @@ -532,8 +535,8 @@ def test_a_relation_lowers_with_the_walk_each_call_takes():
'and its adjoint consumes the value column and produces the key column'
)
assert isinstance(priced, At)
assert (priced.over, priced.into) == (('generator',), ('zone',)), (
'an at produces the fine dims and consumes the coarse'
assert (priced.over, priced.into, priced.joined) == (('generator',), ('zone',), ('snapshot',)), (
'an at produces the fine dims, consumes the coarse, and joins on the rest of the key'
)
p_where = program.variable('p').where
assert p_where is not None
Expand Down
Loading