Skip to content
Closed
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
12 changes: 8 additions & 4 deletions docs/about/architecture.md
Original file line number Diff line number Diff line change
Expand Up @@ -269,10 +269,14 @@ rule. A cost phrased as a rule makes one implementation's choice load-bearing in
the language's rulebook.

0. **The layers are ordered, and imports prove it.** Every module imports only
downward, at module level, with **no exception at all**.
`DELIBERATE_LAZY_IMPORTS` in `tests/test_architecture.py` is empty, and an
undeclared in-function import fails the build. A lazy import here is a cycle
to remove, not to defer.
downward, at module level, and every exception is declared in
`DELIBERATE_LAZY_IMPORTS` in `tests/test_architecture.py` with the cycle it
breaks. An undeclared in-function import fails the build, and a declared one
that disappears fails it too. There are two, one per lane, and they are one
cycle: a `where` may compare arithmetic over parameters, so a mask reads an
expression, while a `cases` expression reads a mask. That recursion is the
language's own grammar, so no ordering of a lane's two modules removes it.
Every other lazy import is a cycle to remove, not to defer.
1. **Core AST is the whole language, and the language is upstream.** Both lanes
consume only core AST. Macros and `piecewise:` are expanded away before
dispatch, and so is a named expression unless it states `cases:`. That one
Expand Down
9 changes: 8 additions & 1 deletion src/lpspec/linopy/where.py
Original file line number Diff line number Diff line change
Expand Up @@ -99,6 +99,13 @@ def evaluate(child: program.WhereNode) -> xr.DataArray:
if isinstance(node, program.VariableDefinedNode):
return absence.present(ctx.model, node.name)

if isinstance(node, program.ExpressionComparisonNode):
from lpspec.linopy.builder import _eval # mask ↔ expression recursion

left, right = _eval(node.left, ctx), _eval(node.right, ctx)
compared = xr.DataArray(_PREDICATE_OPS[node.op](left, right))
return compared.fillna(value=False).astype(bool)

if isinstance(node, (program.ParameterComparisonNode, program.DimensionComparisonNode)):
if isinstance(node, program.ParameterComparisonNode):
arr = dataset[node.name]
Expand Down Expand Up @@ -147,7 +154,7 @@ def evaluate(child: program.WhereNode) -> xr.DataArray:
if isinstance(node, program.OrNode):
return evaluate(node.left) | evaluate(node.right)

assert_never(node)
assert_never(node) # pyrefly: ignore[bad-argument-type] — ArithmeticComparisonNode is in the union and lowering always replaces it (NEVER_LOWERED)


def _defined(arr: xr.DataArray, dtype: str) -> xr.DataArray:
Expand Down
32 changes: 30 additions & 2 deletions src/lpspec/relational/engines/polars/predicates.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,8 @@
built first and the frame read after.

A closed vocabulary of its own — comparisons against a parameter, a dimension
label, a position along a dimension, a relation, and the three connectives. It
label, a position along a dimension, a relation, arithmetic over parameters,
and the three connectives. It
takes the :class:`~lpspec.relational.engines.polars.scope.Scope` as an
argument and holds nothing. :func:`masked` is the product a declaration is
instantiated over, cut by its mask: the one place the two meet.
Expand All @@ -24,6 +25,7 @@
from math_spec import program

from lpspec.errors import DataError, position_out_of_range_message, short_groups_message
from lpspec.relational.engines.polars.fragments import join_on
from lpspec.relational.engines.polars.relations import GROUP_RANK, GROUP_SIZE, Grouping

if TYPE_CHECKING:
Expand Down Expand Up @@ -204,7 +206,33 @@ def join_relation(relation: str, dims: tuple[str, ...], column: str | None) -> s
),
)

def _side(expression: program.ExpressionNode, dims: tuple[str, ...]) -> str:
"""*expression* joined onto the carrier as one value column, and that column's name.

A side of a comparison is arithmetic over parameters — the language
refuses a variable and a ``dual()`` there — so it compiles to constant
fragments alone, and they are added over the comparison's dims the way
a constant side of a constraint is. A coordinate no fragment reaches
stays null so the comparison reads false, which is what every other
atom over a missing value does.
"""
from lpspec.relational.engines.polars.compiler import PolarsCompiler # mask ↔ expression recursion

def attach(frame: pl.LazyFrame, alias: str) -> pl.LazyFrame:
compiled = PolarsCompiler(scope).expression(expression, 'a where clause')
assert not compiled.terms and not compiled.quads, (
'a where side is arithmetic over parameters, which compiles to constants alone'
)
product = frame.select(*dims).unique() if dims else frame.select(pl.lit(0).alias('__one__')).head(1)
added = PolarsCompiler(scope).added(compiled.consts, product, fill=False)
return join_on(frame, added.rename({'cval': alias}), dims, 'left')

return carrier.once(f'__where expression {expression!r}__', attach)

def walk(p: program.WhereNode) -> pl.Expr:
if isinstance(p, program.ExpressionComparisonNode):
left, right = (_side(p.left, p.dims), _side(p.right, p.dims))
return falsy_if_null(_COLUMN_COMPARISONS[p.op](pl.col(left), pl.col(right)))
if isinstance(p, program.ParameterComparisonNode):
return _compare(pl.col(join_param(p.name)), p.op, p.value)
if isinstance(p, program.DimensionComparisonNode):
Expand Down Expand Up @@ -249,7 +277,7 @@ def walk(p: program.WhereNode) -> pl.Expr:
return walk(p.left) | walk(p.right)
if isinstance(p, program.NotNode):
return ~falsy_if_null(walk(p.operand))
assert_never(p)
assert_never(p) # pyrefly: ignore[bad-argument-type] — ArithmeticComparisonNode is in the union and lowering always replaces it (NEVER_LOWERED)

condition = walk(mask.root)
return carrier.frame, condition
Expand Down
65 changes: 60 additions & 5 deletions tests/test_architecture.py
Original file line number Diff line number Diff line change
Expand Up @@ -696,6 +696,46 @@ def test_no_sink_reaches_a_sibling():
)


#: Nodes in the language's unions that ``lower_program`` always replaces, so
#: no lowered plan holds one and neither lane can dispatch on it.
#: ``ArithmeticComparisonNode`` is the resolved form of a where comparison
#: between expressions; lowering rewrites every one into an
#: ``ExpressionComparisonNode`` over program expressions.
NEVER_LOWERED = frozenset({'ArithmeticComparisonNode'})


def test_no_never_lowered_node_survives_lowering():
"""``NEVER_LOWERED`` is a claim about upstream, so it is checked rather than trusted.

An entry that lowering stopped replacing would silently excuse both lanes
from a node they now have to build.
"""
import dataclasses

from math_spec import to_program

from tests.fixtures import DISPATCH_SPEC, override

seen: set[str] = set()

def walk(node):
if node is None:
return
seen.add(type(node).__name__)
for field in dataclasses.fields(node):
child = getattr(node, field.name)
if dataclasses.is_dataclass(child):
walk(child)

for where in ('p_max > cost', '0.5 * p_max > 0', 'NOT p_max > cost'):
mask = to_program(override(DISPATCH_SPEC, **{'variables.p.where': where})).variables['p'].where
walk(None if mask is None else mask.root)
assert seen & NEVER_LOWERED == set(), (
f'{sorted(seen & NEVER_LOWERED)} reached a lowered mask — it is no longer excused from either lane'
)
assert 'ExpressionComparisonNode' in seen, 'the probe must reach the node lowering replaces them with'


def test_every_plan_node_is_handled_by_the_compiler():
"""Two-tier economy: a primitive is not done until the engine consumes it.

Expand Down Expand Up @@ -724,7 +764,11 @@ def test_every_plan_node_is_handled_by_the_compiler():
]
for qualifier, union, module in walkers:
source = module.read_text()
unhandled = [c.__name__ for c in get_args(union) if f'{qualifier}.{c.__name__}' not in source]
unhandled = [
c.__name__
for c in get_args(union)
if c.__name__ not in NEVER_LOWERED and f'{qualifier}.{c.__name__}' not in source
]
assert not unhandled, f'{qualifier} nodes unknown to {module.name}: {unhandled}'


Expand Down Expand Up @@ -991,7 +1035,9 @@ class and cannot collide.

from math_spec import program

declared = {node.__name__ for union in (program.ExpressionNode, program.WhereNode) for node in get_args(union)}
declared = {
node.__name__ for union in (program.ExpressionNode, program.WhereNode) for node in get_args(union)
} - NEVER_LOWERED
assert declared, 'no plan node classes found — the census has nothing to run over'

def dispatched_on(*paths: Path) -> set[str]:
Expand Down Expand Up @@ -1056,9 +1102,18 @@ def test_every_module_is_documented_somewhere():


#: Every in-function ``lpspec`` import in the package, with the cycle it
#: breaks. Empty, and that is the claim: the layers are ordered with no
#: exception at all, so a lazy import is only ever a leftover.
DELIBERATE_LAZY_IMPORTS: dict[tuple[str, str], str] = {}
#: breaks. The two here are one cycle, once per lane: a ``where`` may compare
#: arithmetic over parameters, so a mask reads an expression, and a ``cases``
#: expression reads a mask. The recursion is the language's own grammar, so no
#: ordering of these two modules removes it and each lane's mask walk reaches
#: its expression evaluator inside the one handler that needs it.
DELIBERATE_LAZY_IMPORTS: dict[tuple[str, str], str] = {
('linopy/where.py', 'lpspec.linopy.builder'): 'a where side is an expression, and a cases expression holds a mask',
(
'relational/engines/polars/predicates.py',
'lpspec.relational.engines.polars.compiler',
): 'a where side is an expression, and a cases expression holds a mask',
}


def test_lazy_intra_package_imports_are_all_declared():
Expand Down
22 changes: 22 additions & 0 deletions tests/test_compiler.py
Original file line number Diff line number Diff line change
Expand Up @@ -491,3 +491,25 @@ def test_a_zero_edge_writes_its_rows_like_any_other_fill():
bare = program.Translate(program.Parameter('load'), 'snapshot', 1, wrap=False, fill=None)
vacated = q.expression(bare, 'test').consts[0].frame.collect()
assert vacated['snapshot'].to_list() == [1, 2], 'a bare shift vacates, and that is a gap on purpose'


def test_a_where_side_holding_a_variable_is_caught_rather_than_compiled_as_a_constant():
"""The guard the suite survives deleting, because the language gets there first.

`lower_program` only builds an `ExpressionComparisonNode` from a side it
has already refused a variable and a `dual()` on, so no model reaches
`_side` with a variable term — the guard is the invariant that stays true
if that refusal ever moves. Handed the node directly, it says which
assumption broke instead of folding a variable's coefficient into a value
column and masking on it.
"""
from lpspec.relational.engines.polars.predicates import compile_predicate

mask = program.Mask(
program.ExpressionComparisonNode(program.Variable('p'), '>', program.Constant(0.0), ('snapshot', 'generator'))
)
frame = pl.LazyFrame(schema={'snapshot': pl.Int64, 'generator': pl.String})

with pytest.raises(AssertionError, match='arithmetic over parameters, which compiles to constants alone'):
carrier, _ = compile_predicate(compiler().scope, frame, mask, ('snapshot', 'generator'))
carrier.frame.collect()
18 changes: 17 additions & 1 deletion tests/test_resolution_parity.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,7 +25,6 @@
('where', 'match'),
[
pytest.param('typo_name > 0', "'typo_name' not found", id='a-name-nothing-declares'),
pytest.param('p_max > cost', 'compares two parameters', id='two-parameters-compared'),
pytest.param('generator == snapshot', 'compares against dimension', id='a-dimension-on-the-right'),
pytest.param('nonexistent', "'nonexistent' not found", id='a-bare-name-nothing-declares'),
pytest.param('snapshot', 'bare dimension name is true at every coordinate', id='a-bare-dimension-name'),
Expand Down Expand Up @@ -79,6 +78,18 @@ def test_both_lanes_refuse_a_comparison_that_carries_no_variable(tmp_path, dispa
#: The one position a literal survives to: alone, and false. `True` alone
#: is no mask at all and arrives as `None`.
'False',
#: Two parameters compared, which math-spec 0.0.0-alpha.105 still refused
#: with "compares two parameters, which is not in the language".
'p_max > cost',
#: Arithmetic on a side, and a reduction on one — both sides of a
#: comparison are expressions now, so a side carries what a constant side
#: of a constraint carries.
'0.5 * p_max > 0',
'sum(p_max, over=generator) > 0',
#: Both sides constant, so the comparison is one scalar against another and
#: carries no dimension. The eager lane read the result as an array and got
#: a plain bool, which has no `fillna`.
'1 > 0',
]

#: Predicates this sweep cannot host, with where they are checked instead. The
Expand All @@ -96,6 +107,11 @@ def test_both_lanes_refuse_a_comparison_that_carries_no_variable(tmp_path, dispa
'RelationComparisonNode': 'tests/test_label_coords.py::test_a_where_reads_a_relation',
'RelationPairComparisonNode': 'tests/test_label_coords.py::test_a_relation_where_agrees_with_the_oracle',
'RelationDefinedNode': 'tests/test_label_coords.py::test_a_where_reads_a_relation',
#: The resolved form of an arithmetic comparison, which `lower_program`
#: replaces with `ExpressionComparisonNode` over program expressions. It is
#: in `program.WhereNode` but reaches no lowered mask, so no lane dispatches
#: on it and the sweep above cannot host it.
'ArithmeticComparisonNode': 'math-spec tests/test_lowering.py — never lowered into a program mask',
}


Expand Down
Loading