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
6 changes: 4 additions & 2 deletions src/math_spec/exclusivity.py
Original file line number Diff line number Diff line change
Expand Up @@ -220,8 +220,10 @@ def _subject_of(node: TypedPredicateNode) -> Subject:
return Subject('variable', name)
case DimensionComparisonNode(name=name):
return Subject('dim', name)
case DimensionPositionNode(name=name, by=by, group=group):
return Subject('rank', name, by, group)
case DimensionPositionNode(name=name, partition=partition):
if partition is None:
return Subject('rank', name)
return Subject('rank', name, partition.name, partition.produced)
case LookupDefinedNode(name=name) | LookupComparisonNode(name=name):
return Subject('lookup', name)
case LookupPairComparisonNode(name=name, other=other):
Expand Down
17 changes: 7 additions & 10 deletions src/math_spec/program.py
Original file line number Diff line number Diff line change
Expand Up @@ -1149,20 +1149,17 @@ class DimensionComparisonNode:
class DimensionPositionNode:
"""Compare where a row sits along a dimension against a position — ``position(snapshot) == 0``.

Both sides are integers, negative counting from the end. With ``by`` the
position is counted within each group the lookup makes: ``walked`` is its
key column over ``name``, ``group`` the value columns the group is made
of, and ``dims`` the dimensions of its other key columns, which the frame
carries.
Both sides are integers, negative counting from the end. With a
``partition`` the position is counted within each group the lookup makes,
walked as :class:`Translate` walks one: its consumed column is the key
column over ``name``, the group is its produced columns, and its joined
columns are the other key columns, whose dimensions the frame carries.
"""

name: str
op: PredicateOperator
position: int
by: str | None = None
walked: str | None = None
group: tuple[str, ...] = ()
dims: tuple[str, ...] = ()
partition: Walk | None = None


@dataclass(frozen=True)
Expand Down Expand Up @@ -1316,7 +1313,7 @@ def _atom_dims(atom: TypedPredicateNode) -> frozenset[str]:
case DimensionComparisonNode():
return frozenset({atom.name})
case DimensionPositionNode():
return frozenset({atom.name, *atom.dims})
return frozenset({atom.name, *(atom.partition.joined_dims if atom.partition is not None else ())})
case LookupComparisonNode() | LookupPairComparisonNode() | LookupDefinedNode():
return frozenset(atom.dims)
case _:
Expand Down
14 changes: 10 additions & 4 deletions src/math_spec/resolution.py
Original file line number Diff line number Diff line change
Expand Up @@ -690,6 +690,15 @@ def _walk(
f'{context}: {call}: from= and into= both name {both}, and a walk goes between two sets of columns.'
)
return None
for kwarg, roles in (('from', from_roles), ('into', into_roles)):
dims = [shape.dim(r) for r in roles]
if shared := sorted({d for d in dims if dims.count(d) > 1}):
self.errors.append(
f'{context}: {call}: {kwarg}={list(roles)} names two columns over {shared}, and the operand '
f'carries each dimension once, so nothing says which column its coordinate is read at. Walk '
f'between columns over distinct dimensions.'
)
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)
if not forward and not walk.is_function_read:
Expand Down Expand Up @@ -881,10 +890,7 @@ def _position(self, node: UnresolvedPositionNode) -> DimensionPositionNode | Unr
walk = self._partition_walk(node.by, 'position', node.dimension, node.into)
if walk is None:
return node
(walked,) = walk.consumed
return DimensionPositionNode(
node.dimension, node.op, node.position, node.by, walked, walk.produced, walk.joined_dims
)
return DimensionPositionNode(node.dimension, node.op, node.position, walk)

def _comparison(self, node: UnresolvedComparisonNode) -> WhereNode | UnresolvedWhereNode:
"""``name <op> literal``, or the one structural form ``lookup <op> lookup``."""
Expand Down
11 changes: 6 additions & 5 deletions src/math_spec/typesetting/walk.py
Original file line number Diff line number Diff line change
Expand Up @@ -297,10 +297,11 @@ def _value_read(self, name: str, column: str, ctx: _Context) -> str:

def _position_group(self, node: DimensionPositionNode, ctx: _Context) -> str:
"""The group a grouped position counts within: the lookup's group columns read at the row's key."""
assert node.by is not None
lk = self.schema.lookups[node.by]
keyed = self.format.joined([ctx.subscript(dict(lk.columns)[k]) for k in lk.keys], '')
reads = [self.format.apply(self._column(node.by, column, len(lk.values) == 1), keyed) for column in node.group]
assert node.partition is not None
walk = node.partition
keyed = self.format.joined([ctx.subscript(walk.dim(k)) for k in walk.key], '')
single = len(walk.values) == 1
reads = [self.format.apply(self._column(walk.name, column, single), keyed) for column in walk.produced]
return self._tuple(reads)

def _tuple(self, reads: list[str]) -> str:
Expand Down Expand Up @@ -587,7 +588,7 @@ def _where(self, node: WhereNode, ctx: _Context) -> tuple[str, int]:
)

if isinstance(node, DimensionPositionNode):
grouping = None if node.by is None else self._position_group(node, ctx)
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
Expand Down
10 changes: 9 additions & 1 deletion tests/test_validation.py
Original file line number Diff line number Diff line change
Expand Up @@ -455,7 +455,7 @@ def test_it_resolves(self, mask: str, position: int, by: str | None):
assert isinstance(node, DimensionPositionNode)
assert node.name == 'snapshot'
assert node.position == position
assert node.by == by
assert (node.partition.name if node.partition is not None else None) == by

@pytest.mark.parametrize(
('mask', 'fragments'),
Expand Down Expand Up @@ -634,6 +634,14 @@ class TestRulesDecidedWithoutData:
("from= and into= both name ['h']",),
id='a-from-list-overlapping-to',
),
pytest.param(
{
'lookups.lz': {'over': {'g': 'g', 'h0': 'h', 'h1': 'h'}, 'key': 'g'},
'objective': {'expression': 'sum(sum(p, by=lz, from=[h0, h1], into=g))'},
},
("from=['h0', 'h1'] names two columns over ['h'], and the operand carries each dimension once",),
id='a-from-list-naming-two-columns-over-one-dimension',
),
pytest.param(
{'objective': {'expression': 'sum(shift(p, over=g, offset=1, edge=0, by=lk, from=g))'}},
(
Expand Down
Loading