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
29 changes: 29 additions & 0 deletions linopy/expressions.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,6 +25,7 @@
from typing import (
TYPE_CHECKING,
Any,
Literal,
Self,
TypeAlias,
TypeVar,
Expand Down Expand Up @@ -2496,6 +2497,17 @@ def const(self) -> DataArray:
def const(self, value: DataArray) -> None:
self._data = assign_multiindex_safe(self.data, const=value)

@property
def has_terms(self) -> DataArray:
"""Check live variable terms directly from sparse rows."""
csr = self._csr
if csr is None:
return super().has_terms
present = ~np.isnan(csr.const)
return csr.grid.dataarray(
present & (np.diff(csr.csr.indptr) > 0), name="has_terms"
)

def _combined_with_constant(
self,
self_const: DataArray,
Expand Down Expand Up @@ -2986,6 +2998,23 @@ def rename(
return type(self)._from_csr(csr.renamed(relabel), self._model)
return super().rename(name_dict)

def drop_vars(
self,
names: str | Iterable[Hashable] | Callable[[Dataset], str | Iterable[Hashable]],
*,
errors: Literal["raise", "ignore"] = "raise",
) -> Self:
"""Drop auxiliary coordinates without materializing sparse terms."""
csr = self._csr
if csr is None or callable(names):
return super().drop_vars(names, errors=errors)
names = [names] if isinstance(names, str) else list(names)
if set(names) & (set(csr.grid.dims) | {"coeffs", "vars", "const"}):
return super().drop_vars(names, errors=errors)
coords = csr.grid.to_dataset().drop_vars(names, errors=errors)
grid = Grid.from_dataset(coords, csr.grid.dims)
return type(self)._from_csr(replace(csr, grid=grid), self._model)

def to_quadexpr(self) -> QuadraticExpression:
"""Convert LinearExpression to QuadraticExpression."""
vars = self.data.vars.expand_dims(FACTOR_DIM)
Expand Down
39 changes: 39 additions & 0 deletions test/test_csr.py
Original file line number Diff line number Diff line change
Expand Up @@ -1559,6 +1559,7 @@ def dense_build(build: str) -> LinearExpression:
"coords": lambda e: xr.Dataset(coords=e.coords),
"indexes": lambda e: {k: list(v) for k, v in e.indexes.items()},
"isnull": lambda e: e.isnull(),
"has_terms": lambda e: e.has_terms,
"repr": lambda e: repr(e).splitlines()[2:],
}

Expand Down Expand Up @@ -1588,6 +1589,44 @@ def test_is_sparse_tracks_backing_and_repr_marks_it() -> None:
assert not (1.0 * c.gen_p).is_sparse


@pytest.mark.parametrize("build", list(SPARSE_BUILDS))
def test_drop_aux_coordinates_preserves_sparse_backing(build: str) -> None:
"""
Dropping auxiliary coordinates must preserve terms and sparse storage.

Failures: auxiliary coordinates survive, missing ignored names densify,
the source is mutated, or dropping coordinates changes expression values.
"""
require_v1()
sparse, dense = sparse_and_dense(build)
names = [str(n) for n in sparse.coords if n not in sparse.coord_dims]
names.append("missing-coordinate")
result = sparse.drop_vars(names, errors="ignore")
expected = dense.drop_vars(names, errors="ignore")
assert result.is_sparse
assert sparse.is_sparse
assert not set(names) & set(result.coords)
assert_linequal(result, expected)


def test_drop_missing_coordinate_raises_without_densifying() -> None:
"""A missing coordinate must raise the xarray error without materializing terms."""
require_v1()
sparse, _ = sparse_and_dense("grouped")
with pytest.raises(ValueError, match="cannot be found"):
sparse.drop_vars("missing-coordinate")
assert sparse.is_sparse


def test_has_terms_counts_explicit_zero_coefficients() -> None:
"""Zero coefficients still reference variables and must count as terms."""
require_v1()
c = base_model(sparse=True)
expression = (0.0 * c.gen_p).groupby(c.gbus).sum()
assert expression.has_terms.all().item()
assert expression.is_sparse


def add_chunked(e: LinearExpression, c: Case) -> Any:
chunked = base_model()
chunked.m.chunk = {"bus": 2}
Expand Down
Loading