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 src/mathspec/_expression_resolver.py
Original file line number Diff line number Diff line change
Expand Up @@ -312,7 +312,7 @@ def _at(self, operand: Expression, columns: ColumnsNode) -> Expression | None:
except DimensionError as e:
self.errors.append(str(e))
return None
join = self.lookup(columns, frozenset(axis.dimension for axis in inner))
join = self.lookup(columns, inner)
return None if join is None else Join(operand, join)

def _translation(
Expand Down Expand Up @@ -357,7 +357,7 @@ def _bare_sum(self, operand: Expression) -> Expression | None:
f'scalar. Drop the sum.'
)
return None
return Sum(operand, tuple(sorted(inner, key=str)))
return Sum(operand, tuple(Axis(d) for d in sorted(inner)))

def _edge_fits(self, operand: Expression, offset: int | str, *, wrap: bool, fill: float | None) -> bool:
"""What a ``shift``'s ``edge=`` may say, and where saying nothing is an answer.
Expand Down
11 changes: 5 additions & 6 deletions src/mathspec/_where_resolver.py
Original file line number Diff line number Diff line change
Expand Up @@ -34,13 +34,12 @@
UnresolvedPredicateCallNode,
UnresolvedWhereNode,
)
from mathspec.dimensions import dims_of, frame_of, join_dims
from mathspec.dimensions import dims_of, join_dims
from mathspec.errors import DimensionError, LanguageError, did_you_mean, prefixed
from mathspec.expansion import expand
from mathspec.program import (
Add,
And,
Axis,
BooleanLiteral,
Constant,
CountComparison,
Expand Down Expand Up @@ -232,11 +231,11 @@ def _pulled_back(self, node: UnresolvedPredicateCallNode, mask: Mask) -> Predica
if by is None:
return node
try:
dims = join_dims(by, frame_of(mask.dims), context, 'the predicate')
dims = join_dims(by, mask.dims, context, 'the predicate')
except DimensionError as refusal:
self.errors.append(str(refusal))
return node
return PulledBackPredicate(mask, by, tuple(sorted(axis.dimension for axis in dims)))
return PulledBackPredicate(mask, by, tuple(sorted(dims)))

def _count(self, node: UnresolvedCountNode) -> Predicate | UnresolvedWhereNode:
"""``count(<predicate>, over=<dim>) <op> <integer>`` — how many coordinates the predicate admits.
Expand Down Expand Up @@ -353,7 +352,7 @@ def _expression_comparison(self, node: UnresolvedComparisonNode) -> ExpressionCo
if len(self.errors) > found:
return node
assert len(sides) == 2, 'a side of a where builds or refuses, since a where holds no formal'
dims: set[Axis] = set()
dims: set[str] = set()
for side in sides:
if carries_variable(side):
self.errors.append(
Expand Down Expand Up @@ -381,7 +380,7 @@ def _expression_comparison(self, node: UnresolvedComparisonNode) -> ExpressionCo
f'the comparison.'
)
return node
return ExpressionComparison(left, node.op, right, tuple(d for d in ns.schema.dimensions if Axis(d) in dims))
return ExpressionComparison(left, node.op, right, tuple(d for d in ns.schema.dimensions if d in dims))

def _position(
self, call: FunctionCallNode, node: UnresolvedComparisonNode
Expand Down
2 changes: 1 addition & 1 deletion src/mathspec/composition.py
Original file line number Diff line number Diff line change
Expand Up @@ -502,7 +502,7 @@ def _definer_frame(loaded: Mapping[str, Spec], key: str) -> frozenset[str]:
frame |= set(spec.program.expressions[key].dims)
given = spec.program.given.expressions.get(key)
if given is not None and given.term is not None:
frame |= {str(axis) for axis in dims_of(given.term, spec, f"Given expression '{key}'")}
frame |= dims_of(given.term, spec, f"Given expression '{key}'")
return frozenset(frame)


Expand Down
122 changes: 52 additions & 70 deletions src/mathspec/dimensions.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,7 @@
#
# SPDX-License-Identifier: MIT

"""Static dim-set checking — a type system whose type is a set of axes.
"""Static dim-set checking — a type system whose type is a set of dim names.

Every node's dim set is computable before any data is attached, so this pass runs
at load on the resolved tree. The per-node rules are the "Dim algebra" table in
Expand All @@ -18,7 +18,6 @@
from mathspec.operators import AMOUNTS
from mathspec.program import (
Add,
Axis,
Cases,
Constant,
CountComparison,
Expand Down Expand Up @@ -54,33 +53,12 @@
from mathspec.spec import empty_sums

if TYPE_CHECKING:
from collections.abc import Iterable

from mathspec.program import Program
from mathspec.spec import Spec


def frame_of(dims: Iterable[str]) -> frozenset[Axis]:
"""The axes of a declared frame: each dimension's own, and no join's."""
return frozenset(Axis(d) for d in dims)


def _names(frame: Iterable[Axis]) -> list[str]:
"""A frame as a refusal prints it: each axis as the file writes it, sorted."""
return sorted(str(axis) for axis in frame)


def _dimensions(frame: frozenset[Axis]) -> frozenset[str]:
"""The dimensions whose own axes *frame* holds, which is what a join matches the operand on."""
return frozenset(axis.dimension for axis in frame if axis.column is None)


def dims_of(node: Expression, schema: Spec, context: str) -> frozenset[Axis]:
"""The frame of a resolved expression, checking every rule on the way.

A declaration's own frame holds each dimension's axis. A frame inside a
sum through a relation also holds the axes its join opens, which that sum
takes away.
def dims_of(node: Expression, schema: Spec, context: str) -> frozenset[str]:
"""The dim set of a resolved expression, checking every rule on the way.

Raises:
DimensionError: On the first rule broken.
Expand All @@ -89,14 +67,14 @@ def dims_of(node: Expression, schema: Spec, context: str) -> frozenset[Axis]:
return frozenset()

if isinstance(node, Parameter):
return frame_of({**schema.parameters, **schema.given.parameters}[node.name].dims)
return frozenset({**schema.parameters, **schema.given.parameters}[node.name].dims)

if isinstance(node, Variable):
columns = {**schema.variables, **schema.given.variables, **schema.given.expressions, **empty_sums(schema)}
return frame_of(columns[node.name].dims or ())
return frozenset(columns[node.name].dims or ())

if isinstance(node, Dual):
return frame_of({**schema.constraints, **schema.given.constraints}[node.constraint].dims)
return frozenset({**schema.constraints, **schema.given.constraints}[node.constraint].dims)

if isinstance(node, Named):
return _named_dims(node, schema, context)
Expand All @@ -107,48 +85,56 @@ def dims_of(node: Expression, schema: Spec, context: str) -> frozenset[Axis]:
if isinstance(node, Negate | Add | Multiply | Power | Divide):
return frozenset().union(*(dims_of(child, schema, context) for child in children(node)))

inner = dims_of(node.operand, schema, context)
if isinstance(node, Sum):
return _sum_dims(node, inner, context)
return _sum_dims(node, schema, context)
inner = dims_of(node.operand, schema, context)
if isinstance(node, Join):
assert not node.columns.axes, 'a join that opens axes is read only by the sum that closes them'
return join_dims(node.columns, inner, context, 'the expression')
if isinstance(node, Translate | WindowSum):
return _translation_dims(node, inner, schema, context)

assert_never(node)


def _named_dims(node: Named, schema: Spec, context: str) -> frozenset[Axis]:
def _named_dims(node: Named, schema: Spec, context: str) -> frozenset[str]:
"""An entry's declared frame where it has one — a narrower arm or body broadcasts along the rest — else its body's."""
declared = schema.expressions[node.name].dims
if declared is not None:
return frame_of(declared)
return frozenset(declared)
return dims_of(node.body, schema, context)


def _not_carried(context: str, call: str, inner: frozenset[Axis], rewrite: str) -> str:
def _not_carried(context: str, call: str, inner: frozenset[str], rewrite: str) -> str:
"""The refusal for an operator reaching a dim its operand does not carry; *rewrite* is the operator's own."""
return (
f'{context}: {call} but the expression has dims {_names(inner)}. An operator over a dim the '
f'{context}: {call} but the expression has dims {sorted(inner)}. An operator over a dim the '
f'operand does not carry is a no-op that builds and solves wrong — {rewrite}.'
)


def _sum_dims(node: Sum, inner: frozenset[Axis], context: str) -> frozenset[Axis]:
"""``sum`` reduces each axis in ``over`` away, so the operand carries every one."""
def _sum_dims(node: Sum, schema: Spec, context: str) -> frozenset[str]:
"""``sum`` reduces each dim in ``over`` away, so the operand carries every one.

The axes a join opens never reach a frame: the sum over the join closes
every one, so its frame is the join's frame less them.
"""
if isinstance(node.operand, Join) and node.operand.columns.axes:
join = node.operand
assert node.over == join.columns.axes, 'a sum through a relation closes exactly the axes its join opens'
return join_dims(join.columns, dims_of(join.operand, schema, context), context, 'the expression')
inner = dims_of(node.operand, schema, context)
for summed in node.over:
if summed not in inner:
if summed.dimension not in inner:
raise DimensionError(_not_carried(context, f'sum(over={summed})', inner, 'drop the sum, or fix the dim'))
return inner - frozenset(node.over)
return inner - {axis.dimension for axis in node.over}


def join_dims(columns: JoinColumns, inner: frozenset[Axis], context: str, operand: str) -> frozenset[Axis]:
"""The frame *inner* has once *columns* joins it, an expression's or a predicate's alike.
def join_dims(columns: JoinColumns, inner: frozenset[str], context: str, operand: str) -> frozenset[str]:
"""The dims *inner* has once *columns* joins it, an expression's or a predicate's alike.

The dims joined on go, the dims grouped by arrive, and each column joined
on and not grouped by opens its own axis ([`JoinColumns.axes`][mathspec.program.JoinColumns.axes]), which
the [`Sum`][mathspec.program.Sum] over the join takes away. A column
both joined on and grouped by keeps its dim. The call a refusal quotes is
The dims joined on go and the dims grouped by arrive. A column both
joined on and grouped by keeps its dim. The call a refusal quotes is
``at`` where each group is one row, and ``sum`` otherwise.

Raises:
Expand All @@ -158,39 +144,36 @@ def join_dims(columns: JoinColumns, inner: frozenset[Axis], context: str, operan
lookup = columns.one_row_per_group
named = columns.dropped if lookup else columns.added
call = f'{"at" if lookup else "sum"}(by={columns.name}[{", ".join(named)}])'
carried = _dimensions(inner)
if missing := sorted(set(columns.dropped_dims) - carried):
if missing := sorted(set(columns.dropped_dims) - inner):
if lookup:
raise DimensionError(
f'{context}: {call} joins on {missing}, which {operand} does not carry (dims '
f'{_names(inner)}). A lookup joins the operand on the columns it reads at — '
f'{sorted(inner)}). A lookup joins the operand on the columns it reads at — '
f'sum is the call that groups by them.'
)
raise DimensionError(
_not_carried(context, f'{call} joins on {missing} to sum it away,', inner, 'drop the sum, or fix the dim')
)
added, dropped = set(columns.added_dims), set(columns.dropped_dims)
if clash := sorted((added & carried) - dropped):
if clash := sorted((added & inner) - dropped):
raise DimensionError(
f'{context}: {call} groups by {clash}, which the expression already carries.\n'
f'A join on a column the operand carries matches it rather than grouping by it, so a '
f'call brings the dims it groups by. Move the factor carrying {clash} outside the operator, '
f'or group by a column over another dimension.'
)
_check_joined(call, columns, inner, context)
return (inner - frame_of(columns.joined_dims)) | frame_of(columns.grouped_dims) | frozenset(columns.axes)
return (inner - set(columns.joined_dims)) | set(columns.grouped_dims)


#: The verb a file writes each translation with, which its refusals quote.
_VERBS: dict[type[Translate | WindowSum], str] = {Translate: 'shift', WindowSum: 'sum_back'}


def _translation_dims(
node: Translate | WindowSum, inner: frozenset[Axis], schema: Spec, context: str
) -> frozenset[Axis]:
"""``shift`` and ``sum_back`` keep every axis, and a named amount and a partition are checked here."""
def _translation_dims(node: Translate | WindowSum, inner: frozenset[str], schema: Spec, context: str) -> frozenset[str]:
"""``shift`` and ``sum_back`` keep every dim, and a named amount and a partition are checked here."""
verb = _VERBS[type(node)]
if Axis(node.along) not in inner:
if node.along not in inner:
raise DimensionError(
_not_carried(
context,
Expand All @@ -206,7 +189,7 @@ def _translation_dims(
return inner


def _check_joined(call: str, use: JoinColumns | Partition, inner: frozenset[Axis], 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 All @@ -215,10 +198,10 @@ def _check_joined(call: str, use: JoinColumns | Partition, inner: frozenset[Axis
step along.
"""
dims = use.joined_dims
if missing := sorted(set(dims) - _dimensions(inner)):
if missing := sorted(set(dims) - inner):
raise DimensionError(
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 {_names(inner)}). A join matches "
f"of '{use.name}'), which the expression does not carry (dims {sorted(inner)}). A join matches "
f'the operand on every key column the call does not name — index the operand by them, or '
f'name them in the call.'
)
Expand All @@ -231,7 +214,7 @@ def _check_joined(call: str, use: JoinColumns | Partition, inner: frozenset[Axis


def _check_named_amount(
node: Translate | WindowSum, verb: str, inner: frozenset[Axis], schema: Spec, context: str
node: Translate | WindowSum, verb: str, inner: frozenset[str], schema: Spec, context: str
) -> None:
"""The two rules of an ``offset=`` or ``window=`` naming a parameter that need the operand's dims; resolution holds it to its dtype."""
kwarg, amount = ('offset', node.offset) if isinstance(node, Translate) else ('window', node.width)
Expand All @@ -249,11 +232,11 @@ def _check_named_amount(
groups = (
frozenset(node.partition.dim(v) for v in node.partition.grouped) if node.partition is not None else frozenset()
)
if stray := sorted(frozenset(declared.dims) - _dimensions(inner) - groups):
if stray := sorted(frozenset(declared.dims) - inner - groups):
raise DimensionError(
f'{context}: {verb}({kwarg}={amount}) reads its {words.noun} at the coordinate it '
f"steps from, but '{amount}' varies over {stray}, which that coordinate does not carry "
f'(dims {_names(inner)}). A dim the coordinate does not have is no coordinate at all — '
f'(dims {sorted(inner)}). A dim the coordinate does not have is no coordinate at all — '
f"declare '{amount}' over dims the expression carries, or group by a relation into "
f'one of {stray}, so that each group is reached by its own {words.noun}.'
)
Expand Down Expand Up @@ -303,7 +286,7 @@ def check_schema(schema: Spec, program: Program) -> None:
if given.term is None:
continue
context = f"Given expression '{gname}'"
if extra := _names(dims_of(given.term, schema, context) - frame_of(given.dims)):
if extra := sorted(dims_of(given.term, schema, context) - set(given.dims)):
raise DimensionError(
f'{context}: its term carries {extra}, which its dims {list(given.dims)} do not. A term is read '
f"over the frame the entry states: add {extra} to the entry's dims, or leave them out of the term."
Expand All @@ -314,8 +297,8 @@ def check_schema(schema: Spec, program: Program) -> None:
context = f"Constraint '{cname}'"
_check_where_dims(constraint.where, frame, context)
got = dims_of(constraint.lhs, schema, context) | dims_of(constraint.rhs, schema, context)
if got != frame_of(frame):
stray, missing = _names(got - frame_of(frame)), _names(frame_of(frame) - got)
if got != frame:
stray, missing = sorted(got - frame), sorted(frame - got)
detail = (
f'carries dims {stray} that are not in its dims: {sorted(frame)} — every '
f'stray dim multiplies the rows this constraint builds; add it to '
Expand All @@ -332,7 +315,7 @@ def check_schema(schema: Spec, program: Program) -> None:
got = dims_of(program.objective.expression, schema, context)
if got:
raise DimensionError(
f'{context}: the expression carries dims {_names(got)}, and an objective is one '
f'{context}: the expression carries dims {sorted(got)}, and an objective is one '
f'number. Wrap each additive term in its own sum(): '
f'`sum(p * cost) + sum(p_nom * capex)`.'
)
Expand All @@ -345,22 +328,21 @@ def _check_value_dims(node: Expression, schema: Spec, frame: frozenset[str], con
the second answer a ``dims:`` exists to avoid.
"""
got = dims_of(node, schema, context)
if not got <= frame_of(frame):
if not got <= frame:
raise DimensionError(
f'{context}: the value carries dims {_names(got - frame_of(frame))} outside the dims: '
f'{context}: the value carries dims {sorted(got - frame)} outside the dims: '
f'{sorted(frame)}. A case is a value within the frame — it cannot widen it.'
)


def _check_body_dims(node: Expression, schema: Spec, frame: frozenset[str], context: str) -> None:
"""A plain entry's body may only carry dims its declared frame does; fewer is constant along the rest."""
got = dims_of(node, schema, context)
if not got <= frame_of(frame):
outside = _names(got - frame_of(frame))
if not got <= frame:
raise DimensionError(
f'{context}: the body carries dims {outside} outside the dims: {sorted(frame)}. '
f'{context}: the body carries dims {sorted(got - frame)} outside the dims: {sorted(frame)}. '
f'The dims: are the frame the quantity is read over, and the body cannot widen it: add '
f'{outside} to dims:, or take them out of the body.'
f'{sorted(got - frame)} to dims:, or take them out of the body.'
)


Expand Down
3 changes: 1 addition & 2 deletions src/mathspec/lowering.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,7 +23,6 @@
from mathspec.piecewise import assumptions_of, curve_frame, lp_domain_refusal, resolve_links
from mathspec.program import (
Assumption,
Axis,
BooleanLiteral,
Constant,
ConstraintDeclaration,
Expand Down Expand Up @@ -313,7 +312,7 @@ def _frame_of(name: str, entry: Named, schema: Spec) -> tuple[str, ...]:
if declared is not None:
return tuple(declared)
carried = dims_of(entry.body, schema, f"Named expression '{name}'")
return tuple(d for d in schema.dimensions if Axis(d) in carried)
return tuple(d for d in schema.dimensions if d in carried)


def _bound(value: float | str | None) -> Constant | Parameter | None:
Expand Down
Loading
Loading