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
8 changes: 5 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, WhereNode
from math_spec.program import Direction, Partition, WhereNode

#: 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,8 +133,10 @@ def __str__(self) -> str:

@dataclass(frozen=True)
class RelationNode:
"""A resolved ``by=`` — the relation, and the direction the call reads it in.
"""A resolved ``by=`` — the relation, and the use the call makes of it.

``use`` is a :class:`Direction` for ``sum`` and ``at``, and a
:class:`Partition` for ``shift`` and ``sum_back``.
``dimensions`` is the fine side — what ``sum`` consumes and ``at``
produces — and ``into`` the coarse dims, which ``sum`` produces and ``at``
consumes. The roles joined on are the operand's to carry, and the operator
Expand All @@ -144,7 +146,7 @@ class RelationNode:
name: str
dimensions: tuple[str, ...]
into: tuple[str, ...]
direction: Direction
use: Direction | Partition

def __str__(self) -> str:
return self.name
Expand Down
18 changes: 8 additions & 10 deletions src/math_spec/dimensions.py
Original file line number Diff line number Diff line change
Expand Up @@ -45,6 +45,7 @@
Mask,
ParameterComparisonNode,
ParameterDefinedNode,
Partition,
RelationComparisonNode,
RelationDefinedNode,
RelationPairComparisonNode,
Expand Down Expand Up @@ -228,18 +229,18 @@ def _check_lands_clear(call: str, produced: set[str], consumed: set[str], inner:

def _check_joined(call: str, by: RelationNode, inner: frozenset[str], context: str) -> None:
"""The columns a call joins on are read at their dimensions, so the operand carries every one, each once."""
direction = by.direction
dims = direction.joined_dims
use = by.use
dims = use.joined_dims
if missing := sorted(set(dims) - inner):
raise DimensionError(
f'{context}: {call} joins on {missing} (columns {[r for r in direction.joined if direction.dim(r) in missing]} '
f"of '{direction.name}'), which the expression does not carry (dims {sorted(inner)}). A relation is "
f'{context}: {call} joins on {missing} (columns {[r for r in use.joined if use.dim(r) in missing]} '
f"of '{use.name}'), which the expression does not carry (dims {sorted(inner)}). A relation is "
f'read between two of its columns and joined at the others — index the operand by them, or '
f'read it between different columns.'
)
if twice := sorted({d for d in dims if dims.count(d) > 1 or d in by.dimensions}):
raise DimensionError(
f"{context}: {call} joins '{direction.name}' on {twice} through more than one column, and the operand "
f"{context}: {call} joins '{use.name}' on {twice} through more than one column, and the operand "
f'carries each dimension once. Read between different columns, or use a relation whose joined '
f'columns are over distinct dimensions.'
)
Expand Down Expand Up @@ -421,11 +422,8 @@ def _check_named_amount(node: FunctionCallNode, over: str, inner: frozenset[str]
f"— declare '{amount.name}' over dims '{over}' is not one of."
)
partition = node.kwargs.get('by')
groups = (
frozenset(partition.direction.dim(v) for v in partition.direction.produced)
if isinstance(partition, RelationNode)
else frozenset()
)
use = partition.use if isinstance(partition, RelationNode) else None
groups = frozenset(use.dim(v) for v in use.group) if isinstance(use, Partition) else frozenset()
if stray := sorted(frozenset(declared.dims) - inner - groups):
raise DimensionError(
f'{context}: {node.name}({kwarg}={amount.name}) reads its {words.noun} at the coordinate it '
Expand Down
2 changes: 1 addition & 1 deletion src/math_spec/exclusivity.py
Original file line number Diff line number Diff line change
Expand Up @@ -229,7 +229,7 @@ def _subject_of(node: TypedPredicateNode) -> Subject:
case DimensionPositionNode(name=name, partition=partition):
if partition is None:
return Subject('rank', name)
return Subject('rank', name, partition.name, partition.produced)
return Subject('rank', name, partition.name, partition.group)
case RelationDefinedNode(name=name) | RelationComparisonNode(name=name):
return Subject('relation', name)
case RelationPairComparisonNode(name=name, other=other):
Expand Down
13 changes: 8 additions & 5 deletions src/math_spec/lowering.py
Original file line number Diff line number Diff line change
Expand Up @@ -273,13 +273,15 @@ def sum(self, node: FunctionCallNode) -> program.ExpressionNode:
assert isinstance(consumed, DimensionNode), 'resolution refuses a over= that is not a dimension'
return program.Sum(operand, (consumed.name,))
assert isinstance(by_node, RelationNode), 'resolution refuses a by= that is not a relation'
return program.GroupSum(operand, direction=by_node.direction)
assert isinstance(by_node.use, program.Direction), 'resolution reads sum(by=) in a direction'
return program.GroupSum(operand, direction=by_node.use)

def at(self, node: FunctionCallNode) -> program.ExpressionNode:
"""``at(x, by=relation)`` — the adjoint of :meth:`sum`'s ``by=`` form."""
by_node = node.kwargs['by']
assert isinstance(by_node, RelationNode), 'resolution refuses a by= that is not a relation'
return program.At(self.expr(node.args[0]), direction=by_node.direction)
assert isinstance(by_node.use, program.Direction), 'resolution reads at(by=) in a direction'
return program.At(self.expr(node.args[0]), direction=by_node.use)

def sum_back(self, node: FunctionCallNode) -> program.ExpressionNode:
"""``sum_back(x, along=d, window=w)`` — a trailing window along one dimension.
Expand Down Expand Up @@ -341,8 +343,8 @@ def shift(self, node: FunctionCallNode) -> program.ExpressionNode:
}


def _partition_of(node: FunctionCallNode) -> program.Direction | None:
"""The direction a translation partitions by, if the call names a relation.
def _partition_of(node: FunctionCallNode) -> program.Partition | None:
"""The partition a translation steps inside, if the call names a relation.

That it is a *single* relation, stepped *along the translated dimension*, is
checked with the other dim rules (``math_spec.dimensions``), where a model
Expand All @@ -352,7 +354,8 @@ def _partition_of(node: FunctionCallNode) -> program.Direction | None:
if by_node is None:
return None
assert isinstance(by_node, RelationNode)
return by_node.direction
assert isinstance(by_node.use, program.Partition), "resolution reads a translation's by= as a partition"
return by_node.use


def _bound_expression(value: float | str) -> program.ExpressionNode:
Expand Down
70 changes: 54 additions & 16 deletions src/math_spec/program.py
Original file line number Diff line number Diff line change
Expand Up @@ -81,6 +81,7 @@
'ParameterDeclaration',
'ParameterDefinedNode',
'ParameterDtype',
'Partition',
'PiecewiseDeclaration',
'Power',
'PredicateOperator',
Expand Down Expand Up @@ -342,20 +343,19 @@ class Translate(Expression):
``offset`` is an integer, or the name of an integer parameter that does
not depend on ``dimension`` and carries its sign in the values.

``partition`` is a relation stepped along ``dimension`` — its consumed
column is a key over that dimension, its produced columns are the group —
and the translation then happens inside each group: the neighbour is the
one before in the same group, the edge is the group's, and a wrap closes
each group onto itself. A coordinate the relation sends nowhere reaches
nothing.
``partition`` is a relation with a key column over ``dimension``
(:class:`Partition`), and the translation then happens inside each group
its ``within=`` columns make: the neighbour is the one before in the same
group, the edge is the group's, and a wrap closes each group onto itself.
A coordinate the relation sends nowhere reaches nothing.
"""

operand: ExpressionNode
dimension: str
offset: int | str
wrap: bool
fill: float | None = None
partition: Direction | None = None
partition: Partition | None = None


@dataclass(frozen=True)
Expand Down Expand Up @@ -383,7 +383,7 @@ class Window(Expression):
dimension: str
width: int | str
wrap: bool
partition: Direction | None = None
partition: Partition | None = None


@dataclass(frozen=True)
Expand Down Expand Up @@ -519,9 +519,7 @@ class Direction(NamedTuple):
of ``relation``, which binds every role to its dimension and names the key.
``joined`` is the key roles the call did not name (every role, for a bare
relation): the join keys on them, and a value role left unnamed is not
read. For a partition (``shift``, ``sum_back``, ``position``) ``consumed``
is the key role over the dimension stepped along and ``produced`` the value roles
that make the group, which are the ones ``within=`` named.
read.
"""

relation: RelationDeclaration
Expand Down Expand Up @@ -567,6 +565,48 @@ def is_function_read(self) -> bool:
return bool(self.key) and set(self.key) <= {*self.joined, *self.produced}


class Partition(NamedTuple):
"""One relation as a partition steps along it — the key column stepped along, the group columns, and the key columns joined on.

``along``, ``group`` and ``joined`` are *roles* — column names of
``relation``, which binds every role to its dimension and names the key.
``along`` is the one key column over the dimension stepped along, and
the frame keeps it. ``group`` is the value columns ``within=`` named,
read at the row's key. ``joined`` is the other key columns, whose
dimensions the frame carries. Nothing is consumed and nothing is
produced: the frame does not change.
"""

relation: RelationDeclaration
along: str
group: tuple[str, ...]
joined: tuple[str, ...]

@property
def name(self) -> str:
return self.relation.name

@property
def key(self) -> tuple[str, ...]:
return self.relation.key

@property
def values(self) -> tuple[str, ...]:
return self.relation.values

def dim(self, role: str) -> str:
"""The dimension *role* is bound to."""
return self.relation.dim(role)

@property
def along_dim(self) -> str:
return self.dim(self.along)

@property
def joined_dims(self) -> tuple[str, ...]:
return tuple(self.dim(role) for role in self.joined)


@dataclass(frozen=True)
class DimensionDeclaration:
"""A dimension and the relations with a column over it."""
Expand Down Expand Up @@ -1212,16 +1252,14 @@ 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 a
``partition`` the position is counted within each group the relation makes,
read as :class:`Translate` reads 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.
``partition`` the position is counted within each group the relation makes
(:class:`Partition`), whose joined columns' dimensions the frame carries.
"""

name: str
op: PredicateOperator
position: int
partition: Direction | None = None
partition: Partition | None = None


@dataclass(frozen=True)
Expand Down
41 changes: 20 additions & 21 deletions src/math_spec/resolution.py
Original file line number Diff line number Diff line change
Expand Up @@ -72,6 +72,7 @@
OrNode,
ParameterComparisonNode,
ParameterDefinedNode,
Partition,
PredicateOperator,
RelationComparisonNode,
RelationDeclaration,
Expand Down Expand Up @@ -609,20 +610,18 @@ def _relation_ref(
if 'within' not in named:
return value # the call shape refused it already, with the wording that names the rewrite
over_dim = over.name if isinstance(over, NameNode | DimensionNode) else None
direction = self._partition_direction(name, operator, over_dim, named['within'])
else:
if not ({'over', 'into'} <= set(named)):
return value # the call shape refused it already, with the wording that names the rewrite
direction = self._direction(name, operator, named['over'], named['into'])
partition = self._partition(name, operator, over_dim, named['within'])
if partition is None:
return value
return RelationNode(name, dimensions=(partition.along_dim,), into=(), use=partition)
if not ({'over', 'into'} <= set(named)):
return value # the call shape refused it already, with the wording that names the rewrite
direction = self._direction(name, operator, named['over'], named['into'])
if direction is None:
return value

fine = direction.produced_dims if operator == 'at' else direction.consumed_dims
if operator in ('shift', 'sum_back'):
coarse: tuple[str, ...] = ()
else:
coarse = direction.consumed_dims if operator == 'at' else direction.produced_dims
return RelationNode(name, dimensions=fine, into=coarse, direction=direction)
coarse = direction.consumed_dims if operator == 'at' else direction.produced_dims
return RelationNode(name, dimensions=fine, into=coarse, use=direction)

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 Down Expand Up @@ -708,14 +707,14 @@ def _known_roles(self, name: str, call: str, roles: tuple[str, ...] | None, kwar
return False
return True

def _partition_direction(
def _partition(
self, name: str, operator: str, along_dim: str | None, within_roles: tuple[str, ...]
) -> Direction | None:
"""Which direction a partition (``shift``, ``sum_back``, ``position``) reads relation *name* in along *along_dim*.
) -> Partition | None:
"""How a partition (``shift``, ``sum_back``, ``position``) steps along relation *name* over *along_dim*.

It takes the one key column over that dimension (a key has one column
per dimension), joins on the other key columns and groups by the value
columns *within_roles* names. ``None`` where the dimension is not one
It steps along the one key column over that dimension (a key has one
column per dimension), joins on the other key columns and groups by the
value columns *within_roles* names. ``None`` where the dimension is not one
(already refused), the relation has no key column over it, or
``within=`` names a column that is not a value column.
"""
Expand Down Expand Up @@ -746,7 +745,7 @@ def _partition_direction(
return None
(along,) = over_keys
joined = tuple(r for r in shape.key if r != along)
return Direction(shape, (along,), within_roles, joined)
return Partition(shape, along, within_roles, joined)

def _not_a_relation(self, name: str, operator: str, key: str) -> str | None:
"""Why *name* is not a relation; ``None`` where it is one."""
Expand Down Expand Up @@ -860,10 +859,10 @@ def _position(self, node: UnresolvedPositionNode) -> DimensionPositionNode | Unr
f'are {list(ns.shape_of(node.by).values)}.'
)
return node
direction = self._partition_direction(node.by, 'position', node.dimension, node.into)
if direction is None:
partition = self._partition(node.by, 'position', node.dimension, node.into)
if partition is None:
return node
return DimensionPositionNode(node.dimension, node.op, node.position, direction)
return DimensionPositionNode(node.dimension, node.op, node.position, partition)

def _comparison(self, node: UnresolvedComparisonNode) -> WhereNode | UnresolvedWhereNode:
"""``name <op> literal``, or the one structural form ``relation <op> relation``."""
Expand Down
Loading
Loading