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
2 changes: 1 addition & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -37,7 +37,7 @@ dependencies = [
# A direct reference until math-spec publishes, which is a thing PyPI refuses
# (RELEASING.md) — the same wall the `linopy` extra is behind, with the same
# exit: an ordinary floor the day there is a release to floor against.
"math-spec @ git+https://github.com/energy-models/math-spec@v0.0.0-alpha.89",
"math-spec @ git+https://github.com/energy-models/math-spec@v0.0.0-alpha.92",
# the relational engine: every frame a model is built from, and the parquet
# reader that fills them. Arrow-backed, so a result frame exports the
# PyCapsule protocol and any consumer takes it without this package
Expand Down
12 changes: 1 addition & 11 deletions src/lpspec/linopy/builder.py
Original file line number Diff line number Diff line change
Expand Up @@ -290,7 +290,7 @@ def _eval(node: program.ExpressionNode, ctx: EvaluationContext) -> Any:
_eval(node.operand, ctx),
_relation_arrays(node.coordinate, ctx),
into=node.into,
joined=_joined_dims(node),
joined=node.joined,
labels=ctx.master_coords,
)

Expand Down Expand Up @@ -396,16 +396,6 @@ def _partition(node: program.Translate | program.Window, ctx: EvaluationContext)
return array.rename(node.partition.produced_dims[0])


def _joined_dims(node: program.GroupSum | program.At) -> tuple[str, ...]:
"""The dimensions a node's walks join on — the key columns they neither consume nor produce.

Empty for a map keyed by the one column it is walked out of. A conditioned
map names the rest of its key here, and the operand carries those dims
already, so they are the condition a group is read under.
"""
return tuple(dict.fromkeys(d for walk in node.walks for d in walk.joined_dims))


def _relation_arrays(names: tuple[str, ...], ctx: EvaluationContext) -> tuple[Any, ...]:
"""The declared maps *names* as arrays over the dimensions their keys name, in the order the plan wrote them."""
return tuple(bound_relation(name, ctx.relations) for name in names)
22 changes: 11 additions & 11 deletions src/lpspec/relational/engines/polars/compiler.py
Original file line number Diff line number Diff line change
Expand Up @@ -52,7 +52,7 @@
falsy_if_null,
)
from lpspec.relational.engines.polars.reindex import translate_fragment, window_fragment
from lpspec.relational.engines.polars.relations import joined_dims, landed, mapping, walk_join
from lpspec.relational.engines.polars.relations import landed, mapping, walk_join

if TYPE_CHECKING:
from collections.abc import Callable, Mapping, Sequence
Expand Down Expand Up @@ -669,7 +669,7 @@ def _group_fragment(self, p: TermFragment, g: program.GroupSum, context: str) ->
missing = [d for d in g.over if d not in p.dims]
if missing:
refuse_a_fragment_without_the_dims(p, missing, context, f'sum(by=) over {list(g.over)}')
grouped = self._remap_fragment(p, g.walks)
grouped = self._remap_fragment(p, g)
if p.kind != 'const':
return grouped
return replace(grouped, frame=pl.concat([grouped.frame, self._empty_groups(grouped, g)]))
Expand Down Expand Up @@ -699,8 +699,8 @@ def _empty_groups(self, p: TermFragment, g: program.GroupSum) -> pl.LazyFrame:
spanned = [d for d in p.dims if d not in g.into]
if spanned:
universe = p.frame.select(spanned).unique().join(universe, how='cross')
reached = landed(mapping(self.data.relations, g.walks), g.walks)
empty = universe.join(reached, on=[*joined_dims(g.walks), *g.into], how='anti')
reached = landed(mapping(self.data.relations, g.walks), g)
empty = universe.join(reached, on=[*g.joined, *g.into], how='anti')
return empty.with_columns(pl.lit(0.0, dtype=pl.Float64).alias('cval')).select(*p.dims, *p.carried)

def _at_fragment(self, p: TermFragment, a: program.At, context: str) -> TermFragment:
Expand All @@ -721,7 +721,7 @@ def _at_fragment(self, p: TermFragment, a: program.At, context: str) -> TermFrag
"""
absent = [d for d in a.into if d not in p.dims]
assert not absent, f'in {context}: At through {absent}, which the expression does not span'
remapped = self._remap_fragment(p, a.walks)
remapped = self._remap_fragment(p, a)
return replace(remapped, presences=self._pulled_back_presences(p, a))

def _pulled_back_presences(self, p: TermFragment, a: program.At) -> tuple[Presence, ...]:
Expand All @@ -739,10 +739,10 @@ def _pulled_back_presences(self, p: TermFragment, a: program.At) -> tuple[Presen
dims while this frame keeps the columns that matter — the hazard
:class:`Presence` names.
"""
joined = joined_dims(a.walks)
joined = a.joined
fine = (*joined, *a.over)
table = mapping(self.data.relations, a.walks)
reachable = landed(table, a.walks).unique()
reachable = landed(table, a).unique()
if not p.presences:
total = math.prod(self.data.cardinality[d] for d in fine)
reached = reachable.select(pl.len()).collect().item()
Expand All @@ -756,19 +756,19 @@ def pulled(presence: Presence) -> Presence:
source, keys = (
(presence.frame, keys) if carries_targets else (self.widen(presence.frame, keys, p.dims), p.dims)
)
return Presence(*walk_join(source, table, a.walks, keys))
return Presence(*walk_join(source, table, a, keys))

return tuple(pulled(x) for x in p.presences)

def _remap_fragment(self, p: TermFragment, walks: Sequence[program.Walk]) -> TermFragment:
"""Trade the dims *walks* consume for the ones they produce, through their relations.
def _remap_fragment(self, p: TermFragment, node: program.GroupSum | program.At) -> TermFragment:
"""Trade the dims *node*'s walks consume for the ones they produce, through their relations.

One inner equi-join against :func:`mapping`, keyed as :func:`walk_join`
says. A group consumes the dims its walks are over
(:meth:`_group_fragment`); an ``At`` reads the same tables backwards
(:meth:`_at_fragment`).
"""
frame, dims = walk_join(p.frame, mapping(self.data.relations, walks), walks, p.dims, p.carried)
frame, dims = walk_join(p.frame, mapping(self.data.relations, node.walks), node, p.dims, p.carried)
return TermFragment(dims, frame, p.kind)

def widen(self, presence: pl.LazyFrame, have: tuple[str, ...], want: tuple[str, ...]) -> pl.LazyFrame:
Expand Down
24 changes: 10 additions & 14 deletions src/lpspec/relational/engines/polars/relations.py
Original file line number Diff line number Diff line change
Expand Up @@ -60,11 +60,6 @@ def group_column(role: str) -> str:
# ---------------------------------------------------------------------------


def joined_dims(walks: Sequence[program.Walk]) -> tuple[str, ...]:
"""The dimensions *walks* join on, each once — the key columns neither consumed nor produced, which the operand carries."""
return _each_once(d for walk in walks for d in walk.joined_dims)


def _each_once(dims: Iterable[str]) -> tuple[str, ...]:
return tuple(dict.fromkeys(dims))

Expand Down Expand Up @@ -93,16 +88,16 @@ def _walked(table: pl.LazyFrame, walk: program.Walk) -> pl.LazyFrame:
)


def landed(mapping: pl.LazyFrame, walks: Sequence[program.Walk]) -> pl.LazyFrame:
def landed(mapping: pl.LazyFrame, node: program.GroupSum | program.At) -> pl.LazyFrame:
"""*mapping* at the coordinates it lands on: the joined dimensions and the produced ones, under their names."""
produced = _each_once(d for walk in walks for d in walk.produced_dims)
return mapping.select(*joined_dims(walks), *(pl.col(landing(d)).alias(d) for d in produced))
produced = _each_once(d for walk in node.walks for d in walk.produced_dims)
return mapping.select(*node.joined, *(pl.col(landing(d)).alias(d) for d in produced))


def walk_join(
frame: pl.LazyFrame,
mapping: pl.LazyFrame,
walks: Sequence[program.Walk],
node: program.GroupSum | program.At,
have: Sequence[str],
columns: Sequence[str] = (),
) -> tuple[pl.LazyFrame, tuple[str, ...]]:
Expand All @@ -118,17 +113,18 @@ def walk_join(

Args:
frame: The operand, carrying *have* and *columns*.
mapping: :func:`mapping` for the same walks.
walks: The node's walks.
mapping: :func:`mapping` for the node's walks.
node: The group or the pullback, whose walks say what is consumed
and produced and whose ``joined`` says what is joined on.
have: The dimensions *frame* carries.
columns: The other columns to keep — a fragment's carried ones.

Returns:
The traded frame, and the dimensions it is over, in order.
"""
consumed = _each_once(d for walk in walks for d in walk.consumed_dims)
produced = _each_once(d for walk in walks for d in walk.produced_dims)
joined = joined_dims(walks)
consumed = _each_once(d for walk in node.walks for d in walk.consumed_dims)
produced = _each_once(d for walk in node.walks for d in walk.produced_dims)
joined = node.joined
keep = [d for d in have if d not in consumed]
carried = [d for d in produced if d in keep]
gained = [d for d in produced if d not in keep]
Expand Down
6 changes: 3 additions & 3 deletions uv.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

Loading