diff --git a/doc/release_notes.rst b/doc/release_notes.rst index 0fdb1c0b..4c8bb82f 100644 --- a/doc/release_notes.rst +++ b/doc/release_notes.rst @@ -44,6 +44,10 @@ Upcoming Version * ``linopy.merge`` of several sparse expressions on the same coordinates now adds them in one pass instead of one pairwise sum per operand, and a sparse ``reindex`` that only reorders the existing cells gathers the rows directly instead of rebuilding the matrix. Merging three sparse operands is about 30% faster and such a reindex about 4x faster. (`#1010 `__) * Building ``Model.matrices`` no longer allocates label-sized scaling lookups and skips ``eliminate_zeros`` for frozen constraints. On a model with 2M variables this makes each build of a sparse or frozen model about 20 to 35 ms faster. (`#1008 `__) * Freezing a dense constraint builds the CSR matrix directly from the term rectangle and sorts it only when a row has more than one term. On a model with 2M variables this cuts the freeze of the nodal balance from about 465 to 185 ms and its peak memory from 1.74 GB to 0.25 GB. The export of a mutable constraint also no longer copies a strided term array whose terms are mostly zero. (`#1009 `__) +* ``shift``, ``roll`` and ``diff`` keep a CSR-backed ``LinearExpression`` sparse: the shift is evaluated on the grid's row numbers and the rows are gathered from the sparse backing, so the shifted-in cells are absent as on the dense path, ``roll`` wraps around (``roll_coords`` included) and the constant and auxiliary coordinates follow. ``shift`` with a ``fill_value`` still densifies. On a ``groupby`` result with uneven group sizes (2000 x 720 cells, 800 terms wide) ``shift`` takes 14 ms and 54 MB instead of 1 s and 4.4 GB. (`#1011 `__) +* ``linopy.merge(..., dim=)`` keeps CSR-backed operands sparse: the rows are stacked in input order and the coordinates are what the dense path's ``xr.concat`` yields, so overlapping and duplicate labels along ``dim``, the ``join`` of the other dimensions (cells the join creates are absent), the v1 label check of the auto-detected join and auxiliary coordinates behave as on the dense path. The operands are left sparse, too. A merge along a new dimension, with extra ``xr.concat`` arguments or over non-unique labels that need aligning still densifies, with a notice. (`#1011 `__) +* ``linopy.options["warn_on_densify"]`` now defaults to ``None``, which means: on in a sparse model (``Model(sparse=True)``), off otherwise. Every operation that implicitly drops a sparse backing in a sparse model, e.g. ``.data`` or an operation without a sparse path, emits a ``PerformanceWarning`` naming the reason. Explicit conversions (``CSRConstraint.mutable()``/``to_dense()``, ``add_constraints(..., freeze=False)``) only warn when the option is ``True``. Set the option to ``False`` to silence it, or to ``True`` to see it in a dense model and on explicit conversions as well. (`#969 `__) +* A merge of operands over different dimensions, e.g. ``bal + s`` with ``s`` over a subset of the dimensions of ``bal``, still densifies and now warns in a sparse model. (`#969 `__) *Other* @@ -68,7 +72,7 @@ Upcoming Version * ``densify_terms`` (used by ``sum(drop_zeros=True)`` and the sparse ``@`` path) is now fully vectorised. It previously counted the non-zero positions with a Python loop that scaled quadratically in the number of non-zero terms — 127 s for a (2000 x 60) expression, now 3 ms — and allocated the compacted output at the full original term width. It now allocates only the compacted width and returns the expression unchanged when it holds no zeros. * Under v1, ``@``/``dot`` against a constant now runs as sparse linear algebra instead of building the dense broadcast intermediate (``self * other`` then ``.sum()``): peak memory scales with ``nnz(C) x nterm`` rather than with the full broadcast shape. A ``CSRLinearExpression``-backed operand (from ``groupby(...).sum(sparse=True)``) stays CSR-backed through ``@``, and the result is the compact canonical form (duplicate variables summed, terms label-ordered, explicit zeros pruned — cell activeness is carried by ``const`` alone). (`#748 `__, `#756 `__, `#925 `__) * Reading metadata of a CSR-backed ``LinearExpression`` (``repr()``, ``shape``, ``sizes``, ``dims``, ``coords``, ``indexes``, ``isnull()``) no longer converts it to dense; it is served from the sparse backing, auxiliary coordinates included. (`#962 `__) -* New ``LinearExpression.is_sparse`` tells whether an expression is CSR-backed, and its repr header reads ``LinearExpression (sparse)``. The opt-in option ``linopy.options["warn_on_densify"]`` (default ``False``) emits a ``PerformanceWarning`` naming the reason whenever a sparse backing is dropped, e.g. on ``.data``, an operation without a sparse path, or ``CSRConstraint.mutable()``. (`#969 `__) +* New ``LinearExpression.is_sparse`` tells whether an expression is CSR-backed, and its repr header reads ``LinearExpression (sparse)``. The option ``linopy.options["warn_on_densify"]`` emits a ``PerformanceWarning`` naming the reason whenever a sparse backing is dropped, e.g. on ``.data``, an operation without a sparse path, or ``CSRConstraint.mutable()``. (`#969 `__) * Multiplying, dividing, adding or subtracting a non-scalar constant (``DataArray``, ``Series``, ``ndarray``, …) keeps a CSR-backed ``LinearExpression`` sparse: the coefficients are scaled per cell and only the constant changes, instead of expanding the dense rectangle. The constant is aligned exactly as on the dense path (NaN, label mismatch, joins and auxiliary coordinates); an operand that adds dimensions still densifies. ``expr.const`` is now served from the sparse backing as well. (`#965 `__) * ``sum(dim=...)`` and a further ``groupby(...).sum()`` keep a CSR-backed ``LinearExpression`` sparse: summed rows are merged in the sparse backing instead of stacking the summed dimensions into the term dimension of the dense rectangle. Auxiliary coordinates, absent cells and the constant follow the dense path; ``groupby(...).sum(sparse=False)`` still densifies. (`#964 `__) * ``where``, ``sel``, ``isel``, ``loc`` and ``[]`` keep a CSR-backed ``LinearExpression`` sparse: the selection is evaluated on the grid's row numbers with the dense (xarray) semantics, including ``drop=``, ``method=``, scalar coordinates left by a scalar selection and auxiliary coordinates, and the selected rows are gathered from the sparse backing; masked cells become absent. A condition or indexer that introduces dimensions still densifies. A merge of differing grids now stays sparse when operands carry auxiliary coordinates, which are checked for conflicts as on the dense path. (`#966 `__) diff --git a/linopy/config.py b/linopy/config.py index abc6f242..ccccd77e 100644 --- a/linopy/config.py +++ b/linopy/config.py @@ -96,5 +96,5 @@ def __repr__(self) -> str: display_max_terms=6, semantics=LEGACY_SEMANTICS, sparse_groupby=False, - warn_on_densify=False, + warn_on_densify=None, ) diff --git a/linopy/constraints.py b/linopy/constraints.py index fa42411c..47bddecd 100644 --- a/linopy/constraints.py +++ b/linopy/constraints.py @@ -1520,7 +1520,11 @@ def freeze(self) -> CSRConstraint: def to_dense(self) -> Constraint: """Convert to a Constraint.""" - _densify_notice("frozen constraint converted by `to_dense()`/`mutable()`") + _densify_notice( + "frozen constraint converted by `to_dense()`/`mutable()`", + self._model, + explicit=True, + ) return Constraint(self.data, self._model, self._name) def mutable(self) -> Constraint: @@ -2787,7 +2791,11 @@ def set_blocks(self, block_map: np.ndarray) -> None: for name, constraint in self.items(): if not isinstance(constraint, Constraint): - self.data[name] = constraint = constraint.mutable() + _densify_notice( + "frozen constraint converted by `set_blocks`", self.model + ) + constraint = Constraint(constraint.data, self.model, name) + self.data[name] = constraint res = xr.full_like(constraint.labels, N + 1, dtype=block_map.dtype) entries = replace_by_map(constraint.vars, block_map) diff --git a/linopy/csr.py b/linopy/csr.py index 7a551ba0..7c27463a 100644 --- a/linopy/csr.py +++ b/linopy/csr.py @@ -536,6 +536,44 @@ def taken(self, rows: np.ndarray, grid: Grid) -> CSRLinearExpression: csr = self.csr[np.where(present, rows, 0)] return replace(self, csr=csr, grid=grid).with_const(const) + def concatenated( + self, others: Iterable[CSRLinearExpression], dim: str, grid: Grid + ) -> CSRLinearExpression: + """ + Stack this and ``others``, in that order, along ``dim`` into the + cells of ``grid``, all operands sharing the labels of the other dims + in this grid's dim order. Rows are gathered with their terms, + explicit zeros included, and their constants. Auxiliary coordinates + are ``grid``'s. + """ + parts = (self, *others) + n_cols = max(p.csr.shape[1] for p in parts) + blocks = [ + scipy.sparse.csr_array( + (p.csr.data, p.csr.indices, p.csr.indptr), shape=(p.n_cells, n_cols) + ) + for p in parts + ] + csr = scipy.sparse.vstack(blocks, format="csr") + if csr.shape[0] != grid.size: + raise ValueError( + f"Stacked {csr.shape[0]} rows into a grid of {grid.size} cells." + ) + const = np.concatenate([p.const for p in parts]) + stacked = replace(self, csr=csr, const=const, grid=grid) + axis = self.grid.dims.index(dim) + if not axis: + return stacked + offsets = np.cumsum([0, *(p.n_cells for p in parts[:-1])]) + order = np.concatenate( + [ + (o + np.arange(p.n_cells)).reshape(p.shape) + for o, p in zip(offsets, parts) + ], + axis=axis, + ).reshape(-1) + return stacked.taken(order, grid) + def reindexed(self, grid: Grid, fill: float = np.nan) -> CSRLinearExpression: """ Remap rows onto a new grid, possibly in a new dim order: dropped @@ -549,12 +587,7 @@ def reindexed(self, grid: Grid, fill: float = np.nan) -> CSRLinearExpression: source = np.full(grid.size, -1, dtype=row_map.dtype) source[row_map] = np.arange(grid.size) if (source >= 0).all(): - return replace( - self, - csr=self.csr[source], - const=self.const[source], - grid=self.grid.conformed(grid), - ) + return self.taken(source, self.grid.conformed(grid)) coo = self.csr.tocoo() keep = valid[coo.coords[0]] @@ -781,12 +814,15 @@ def csr_to_term_arrays( return vars_, coeffs -def _densify_notice(reason: str) -> None: +def _densify_notice(reason: str, model: Model, explicit: bool = False) -> None: """ Emit a :class:`~linopy.constants.PerformanceWarning` naming why a sparse - (CSR) backing is dropped, if ``options["warn_on_densify"]`` is set. + (CSR) backing is dropped, if ``options["warn_on_densify"]`` is set, or + left at its default ``None``, ``model`` is sparse and the conversion is + an implicit fallback rather than an ``explicit`` user request. """ - if options["warn_on_densify"]: + enabled = options["warn_on_densify"] + if model.sparse and not explicit if enabled is None else enabled: warn_outside_linopy( f"Sparse (CSR) backing densified: {reason}.", PerformanceWarning ) diff --git a/linopy/expressions.py b/linopy/expressions.py index bc811abd..97f0bf45 100644 --- a/linopy/expressions.py +++ b/linopy/expressions.py @@ -666,7 +666,9 @@ def sum( "existing dimension, without use_fallback." ) if csr is None: - _densify_notice("groupby-sum with a grouper without a sparse path") + _densify_notice( + "groupby-sum with a grouper without a sparse path", self.model + ) if multikey_frame is not None: group = multikey_frame @@ -2609,31 +2611,51 @@ def _selected( flat = rows.fillna(-1).to_numpy().reshape(-1).astype(np.int64) return csr.taken(flat, grid) - def sel(self, *args: Any, **kwargs: Any) -> LinearExpression: + def _gathered( + self, name: str, dense: Callable[..., Self], /, *args: Any, **kwargs: Any + ) -> Self: """ - Select by label as ``Dataset.sel``. For a CSR-backed expression, - returns a CSR-backed result when the selection stays on the grid. + Apply the dense method ``name`` to the grid's row numbers via + :meth:`_selected`, falling back to ``dense`` off the grid. """ - csr = self._selected(lambda rows: rows.sel(*args, **kwargs), "sel") + csr = self._selected(operator.methodcaller(name, *args, **kwargs), name) if csr is None: - return super().sel(*args, **kwargs) + return dense(*args, **kwargs) return type(self)._from_csr(csr, self._model) - def isel(self, *args: Any, **kwargs: Any) -> LinearExpression: + def sel(self, *args: Any, **kwargs: Any) -> Self: + """ + Select by label as ``Dataset.sel``. For a CSR-backed expression, + returns a CSR-backed result when the selection stays on the grid. + """ + return self._gathered("sel", super().sel, *args, **kwargs) + + def isel(self, *args: Any, **kwargs: Any) -> Self: """ Select by position as ``Dataset.isel``. For a CSR-backed expression, returns a CSR-backed result when the selection stays on the grid. """ - csr = self._selected(lambda rows: rows.isel(*args, **kwargs), "isel") - if csr is None: - return super().isel(*args, **kwargs) - return type(self)._from_csr(csr, self._model) + return self._gathered("isel", super().isel, *args, **kwargs) def __getitem__(self, selector: int | tuple[slice, list[int]] | slice) -> Self: - csr = self._selected(lambda rows: rows[selector], "__getitem__") - if csr is None: - return super().__getitem__(selector) - return type(self)._from_csr(csr, self._model) + return self._gathered("__getitem__", super().__getitem__, selector) + + def shift(self, *args: Any, **kwargs: Any) -> Self: + """ + Shift along dimensions as ``Dataset.shift``. For a CSR-backed + expression, returns a CSR-backed result with the shifted-in cells + absent, unless a ``fill_value`` is passed. + """ + if len(args) > 1 or "fill_value" in kwargs: + self._densify("`shift` with a fill_value") + return self._gathered("shift", super().shift, *args, **kwargs) + + def roll(self, *args: Any, **kwargs: Any) -> Self: + """ + Roll along dimensions as ``Dataset.roll``. For a CSR-backed + expression, returns a CSR-backed result. + """ + return self._gathered("roll", super().roll, *args, **kwargs) def where( self, @@ -3703,7 +3725,7 @@ def _densify_all(exprs: Iterable[Any], reason: str) -> None: if isinstance(e, LinearExpression) and (csr := e._csr) is not None ] if sparse: - _densify_notice(reason) + _densify_notice(reason, sparse[0][0].model) for e, csr in sparse: e._data = csr.to_dense()._data e._csr = None @@ -3737,6 +3759,40 @@ def _aligned( return [p.reindexed(grid, fill) for p in csrs] +def _concatenated( + csrs: list[CSRLinearExpression], dim: str, join: JoinOptions | None +) -> CSRLinearExpression | str: + """ + Stack CSR expressions sharing one dim order along the grid dim ``dim``, + in order. The result's coordinates are what the dense path's + ``xr.concat`` yields on the grids alone: labels and auxiliary + coordinates concatenated along ``dim``, the other dims joined as + ``join`` says (``outer`` by default, after the v1 label check for the + auto-detected join), with the cells the join creates absent. Returns the + reason as a string where the dense path owns the semantics instead: a + join over non-unique labels. + """ + dims = csrs[0].grid.dims + metadata = [p.grid.to_dataset() for p in csrs] + if join is None: + enforce_merge_dims(metadata, concat_dim=dim, context=f"merge along dim {dim!r}") + enforce_aux_conflict(metadata, concat_dim=dim) + combined = xr.concat( + metadata, dim, join=join or "outer", coords="minimal", compat="override" + ) + grid = Grid.from_dataset(combined, dims) + aligned = [] + for p in csrs: + target = grid.with_indexes({dim: p.grid.indexes[dim]}) + if join == "override" or p.grid.same_layout(target): + aligned.append(replace(p, grid=target)) + elif p.grid.is_unique: + aligned.append(p.reindexed(target)) + else: + return "merge over non-unique labels" + return aligned[0].concatenated(aligned[1:], dim, grid) + + def _try_csr_merge( exprs: Any, dim: str, @@ -3747,20 +3803,19 @@ def _try_csr_merge( """ Sparse branch of :func:`merge`: combine plain LinearExpressions over one set of grid dimensions (CSR-backed or dense-convertible) as sparse matrix - addition. Grids that share dims in a different order are transposed onto - the template order first. Grids that differ in their labels are aligned - row-wise onto the joined grid, the cells the join creates carrying the - fill of the dense path (zero, or NaN for ``fill_value=ABSENT``). Auxiliary - coordinates are checked for conflicts on the operands as given and follow - their rows onto the joined grid. Returns None to fall through to the - dense path. + addition along the term dimension, or as a row stack along one of the + grid dimensions (:func:`_concatenated`). Grids that share dims in a + different order are transposed onto the template order first. Grids that + differ in their labels are aligned row-wise onto the joined grid, the + cells the join creates carrying the fill of the dense path (zero, or NaN + for ``fill_value=ABSENT``). Auxiliary coordinates are checked for + conflicts on the operands as given and follow their rows onto the joined + grid. Returns None to fall through to the dense path. """ if not any(type(e) is LinearExpression and e._csr is not None for e in exprs): return None - if dim != TERM_DIM or kwargs: - _densify_all( - exprs, "merge along a coordinate dimension or with extra arguments" - ) + if kwargs: + _densify_all(exprs, "merge with extra arguments") return None if not all(type(e) is LinearExpression for e in exprs): _densify_all(exprs, "merge with an operand that is not a LinearExpression") @@ -3769,6 +3824,9 @@ def _try_csr_merge( if any(set(e.coord_dims) != dims for e in exprs[1:]): _densify_all(exprs, "merge of operands over different dimensions") return None + if dim != TERM_DIM and dim not in dims: + _densify_all(exprs, "merge along a new dimension") + return None for e in exprs: if e._csr is None and set(e.data.coords) - dims != set( _aux_coords(e.data, dims) @@ -3785,6 +3843,12 @@ def _try_csr_merge( p.reindexed(p.grid.reordered(order)) if p.grid.dims != order else p for p in csrs ] + if dim != TERM_DIM: + stacked = _concatenated(csrs, dim, join) + if isinstance(stacked, str): + _densify_all(exprs, stacked) + return None + return LinearExpression._from_csr(stacked, exprs[0].model) if not all(template.same_grid(p) for p in csrs[1:]): enforce_aux_conflict([Dataset(coords=p.grid.aux) for p in csrs]) aligned = _aligned(csrs, join, join_fill(fill_value, 0.0)) diff --git a/linopy/model.py b/linopy/model.py index 8c041441..30b42090 100644 --- a/linopy/model.py +++ b/linopy/model.py @@ -1399,7 +1399,7 @@ def add_constraints( reason = "chunked model, `Model.chunk` adds constraints unfrozen" else: reason = "constraint added unfrozen, `freeze=False`" - _densify_notice(reason) + _densify_notice(reason, self, explicit=not chunked) data = con.data _check_infinities(data.sign, data.rhs, name) diff --git a/test/test_csr.py b/test/test_csr.py index baf9258f..d51bba86 100644 --- a/test/test_csr.py +++ b/test/test_csr.py @@ -36,6 +36,10 @@ assert_varequal, ) +pytestmark = pytest.mark.filterwarnings( + "ignore:Sparse \\(CSR\\) backing densified:linopy.PerformanceWarning" +) + def require_v1() -> None: if not is_v1(): @@ -1780,6 +1784,11 @@ def add_chunked(e: LinearExpression, c: Case) -> Any: return chunked.m.add_constraints(lhs >= 1, freeze=True) +def set_blocks(e: LinearExpression, c: Case) -> None: + c.m.add_constraints(e >= 1, name="blocked") + c.m.constraints.set_blocks(np.zeros(c.m._xCounter, dtype=int)) + + DENSIFY_OPS: dict[str, tuple[Callable[[LinearExpression, Case], Any], str]] = { "data": (lambda e, c: e.data, "`.data` read"), "merge": (lambda e, c: e + 1.0 * c.flow, "over different dimensions"), @@ -1795,6 +1804,7 @@ def add_chunked(e: LinearExpression, c: Case) -> Any: "constraint added unfrozen, `freeze=False`", ), "add-chunked": (add_chunked, "chunked model"), + "set-blocks": (set_blocks, "frozen constraint converted by `set_blocks`"), "new-dim": (lambda e, c: e * xr.DataArray([1.0, 2.0], coords=[LOC]), "new dim"), "sum-kwargs": (lambda e, c: e.sum(dims="bus"), "`.data` read"), "groupby-fallback": ( @@ -1813,6 +1823,23 @@ def add_chunked(e: LinearExpression, c: Case) -> Any: lambda e, c: e.where(lambda ds: ds.const > 0), "`where` with a condition that is no DataArray", ), + "shift-fill": (lambda e, c: e.shift(bus=1, fill_value=0), "`shift` with a fill"), + "concat-new-dim": ( + lambda e, c: linopy.merge([e, e], dim="new"), + "merge along a new dimension", + ), + "concat-kwargs": ( + lambda e, c: linopy.merge([e, e], dim="snapshot", coords="minimal"), + "merge with extra arguments", + ), + "concat-non-unique": ( + lambda e, c: linopy.merge( + [e.isel(snapshot=[0, 0], bus=[0, 2]), e.isel(snapshot=[1])], + dim="snapshot", + join="outer", + ), + "merge over non-unique labels", + ), } @@ -1822,9 +1849,11 @@ def halves(e: LinearExpression) -> pd.Series: return pd.Series(np.arange(len(index)) % 2, index=index).map({0: "a", 1: "b"}) -@pytest.mark.parametrize("enabled", [True, False], ids=["enabled", "default"]) +@pytest.mark.parametrize( + "enabled", [True, None, False], ids=["enabled", "default", "disabled"] +) @pytest.mark.parametrize("op", list(DENSIFY_OPS)) -def test_warn_on_densify_names_the_reason(op: str, enabled: bool) -> None: +def test_warn_on_densify_names_the_reason(op: str, enabled: bool | None) -> None: require_v1() c = base_model(sparse=True) sparse = SPARSE_BUILDS["grouped"](c) @@ -1834,7 +1863,8 @@ def test_warn_on_densify_names_the_reason(op: str, enabled: bool) -> None: opts.set_value(warn_on_densify=enabled) func(sparse, c) notices = [w for w in caught if issubclass(w.category, linopy.PerformanceWarning)] - if not enabled: + implicit_in_sparse_model = op not in {"add-chunked", "mutable", "add-unfrozen"} + if enabled is False or (enabled is None and not implicit_in_sparse_model): assert notices == [] return assert len(notices) == 1 @@ -1842,6 +1872,18 @@ def test_warn_on_densify_names_the_reason(op: str, enabled: bool) -> None: assert notices[0].filename == __file__ +def test_warn_on_densify_default_is_quiet_in_dense_model() -> None: + require_v1() + c = base_model() + with pytest.warns(FutureWarning, match="deprecated"): + sparse = (c.eff * c.gen_p).groupby(c.gbus).sum(sparse=True) + assert sparse.is_sparse + with warnings.catch_warnings(record=True) as caught: + warnings.simplefilter("always", linopy.PerformanceWarning) + sparse.data + assert not any(issubclass(w.category, linopy.PerformanceWarning) for w in caught) + + @pytest.mark.parametrize("scale", [1.0, 0.0], ids=["plain", "zeros"]) @pytest.mark.parametrize("build", list(SPARSE_BUILDS)) def test_flat_and_to_polars_served_without_densifying(build: str, scale: float) -> None: @@ -2788,3 +2830,179 @@ def test_copy_frozen_matrices(deep: bool) -> None: with pytest.raises(ValueError, match="read-only"): c.constraints["c"].coords["aux"].values[0] = -1 assert c.constraints["c"].coords["aux"].equals(m.constraints["c"].coords["aux"]) + + +SHIFT_OPS: dict[str, Callable[[LinearExpression], LinearExpression]] = { + "shift": lambda e: e.shift({dim0(e): 1}), + "shift-back": lambda e: e.shift({dim0(e): -2}), + "shift-all": lambda e: e.shift({str(d): 1 for d in e.coord_dims}), + "roll": lambda e: e.roll({dim0(e): 1}), + "roll-coords": lambda e: e.roll({dim0(e): -1}, roll_coords=True), + "diff": lambda e: e.diff(dim0(e)), + "diff-2": lambda e: e.diff(dim0(e), 2), +} + + +@pytest.mark.parametrize("scale", [1.0, 0.0], ids=["plain", "zeros"]) +@pytest.mark.parametrize("op", list(SHIFT_OPS)) +@pytest.mark.parametrize("build", ["grouped", "aux", "absent"]) +def test_shift_roll_diff_stay_csr_and_match_dense( + build: str, op: str, scale: float +) -> None: + require_v1() + sparse, dense = sparse_and_dense(build) + sparse, dense = scale * sparse, scale * dense + with no_densify(): + res = SHIFT_OPS[op](sparse) + assert_sparse_matches(res, SHIFT_OPS[op](dense)) + + +@pytest.mark.parametrize( + "op", + [lambda e: e.shift({dim0(e): 1}, 0), lambda e: e.roll({dim0(e): 1}, True)], + ids=["shift-fill", "roll-coords"], +) +def test_shift_roll_take_positional_args( + op: Callable[[LinearExpression], LinearExpression], +) -> None: + require_v1() + sparse, dense = sparse_and_dense("grouped") + assert_linequal(op(sparse), op(dense)) + + +def concat( + parts: list[LinearExpression], + join: JoinOptions | None = None, + dim: str = "snapshot", +) -> LinearExpression: + return linopy.merge(parts, dim=dim, join=join, cls=LinearExpression) + + +def split(e: LinearExpression, dim: str, *parts: list[int]) -> list[LinearExpression]: + return [e.isel({dim: p}) for p in parts] + + +CONCAT_PARTS: dict[str, Callable[[int], list[list[int]]]] = { + "halves": lambda n: [list(range(n // 2)), list(range(n // 2, n))], + "three": lambda n: [[0], list(range(1, n - 1)), [n - 1]], + "overlap": lambda n: [[0, 1], list(range(1, n))], + "reordered": lambda n: [[n - 1, 1], [0, *range(2, n - 1)]], + "duplicate-within": lambda n: [[0, 0], [1]], + "empty-part": lambda n: [[], list(range(1, n))], + "single": lambda n: [[n - 1, 0]], +} + + +@pytest.mark.parametrize("scale", [1.0, 0.0], ids=["plain", "zeros"]) +@pytest.mark.parametrize("parts", list(CONCAT_PARTS)) +@pytest.mark.parametrize("axis", ["first", "last"]) +@pytest.mark.parametrize("build", ["grouped", "aux", "absent"]) +def test_concat_stays_csr_and_matches_dense( + build: str, axis: str, parts: str, scale: float +) -> None: + require_v1() + sparse, dense = sparse_and_dense(build) + sparse, dense = scale * sparse, scale * dense + dim = str(sparse.coord_dims[0 if axis == "first" else -1]) + positions = CONCAT_PARTS[parts](sparse.sizes[dim]) + operands = split(sparse, dim, *positions) + with no_densify(): + res = concat(operands, dim=dim) + assert all(e.is_sparse for e in operands) + assert_sparse_matches(res, concat(split(dense, dim, *positions), dim=dim)) + + +@pytest.mark.parametrize("transposed", [False, True], ids=["same-order", "transposed"]) +def test_concat_with_dense_operand_stays_csr_and_matches_dense( + transposed: bool, +) -> None: + require_v1() + sparse, dense = sparse_and_dense("grouped") + tail = dense.isel(snapshot=[1, 2]) + if transposed: + tail = LinearExpression( + tail.data.transpose(*reversed(tail.coord_dims), ...), tail.model + ) + with no_densify(): + res = concat([sparse.isel(snapshot=[0]), tail]) + assert_sparse_matches(res, concat([dense.isel(snapshot=[0]), tail])) + + +def joined_parts(e: LinearExpression) -> list[LinearExpression]: + """Two snapshot blocks, the second over a subset of the buses.""" + return [e.isel(snapshot=[0]), e.isel(snapshot=[1, 2], bus=[0, 2])] + + +@pytest.mark.parametrize("join", ["outer", "inner", "left", "right"]) +def test_concat_join_stays_csr_and_matches_dense(join: JoinOptions) -> None: + require_v1() + sparse, dense = sparse_and_dense("grouped") + with no_densify(): + res = concat(joined_parts(sparse), join) + assert_sparse_matches(res, concat(joined_parts(dense), join)) + + +def override_parts(e: LinearExpression) -> list[LinearExpression]: + buses = list(e.indexes["bus"]) + return [e.isel(snapshot=[0]), e.isel(snapshot=[1]).reindex(bus=buses[1:] + ["new"])] + + +def one_sided_aux_parts(e: LinearExpression) -> list[LinearExpression]: + return [e.isel(snapshot=[0]), e.isel(snapshot=[1, 2]).drop_vars(["bus", "tag"])] + + +@pytest.mark.parametrize( + ("build", "parts", "join"), + [ + ("grouped", override_parts, "override"), + ("aux", one_sided_aux_parts, "outer"), + ("aux", one_sided_aux_parts, None), + ], +) +def test_concat_coords_follow_dense( + build: str, + parts: Callable[[LinearExpression], list[LinearExpression]], + join: JoinOptions | None, +) -> None: + require_v1() + sparse, dense = sparse_and_dense(build) + with no_densify(): + res = concat(parts(sparse), join) + assert_sparse_matches(res, concat(parts(dense), join)) + + +def test_concat_freezes_like_dense() -> None: + require_v1() + dense_c, sparse_c = twin_models() + + def lhs(c: Case) -> LinearExpression: + parts = split(c.gen_sum(), "bus", [0, 1], [2, 3, 4]) + return linopy.merge(parts, dim="bus", cls=LinearExpression) + + with no_densify(): + frozen = sparse_c.m.add_constraints(lhs(sparse_c) >= 1, name="c") + assert isinstance(frozen, CSRConstraint) + assert_frozen_equal(dense_c.m.add_constraints(lhs(dense_c) >= 1, name="c"), frozen) + + +@pytest.mark.parametrize( + ("parts", "join"), + [ + (joined_parts, None), + (joined_parts, "exact"), + (joined_parts, "override"), + (lambda e: [e.isel(snapshot=[0]), e.isel(snapshot=[1], group=[0])], "outer"), + ], + ids=["auto-mismatch", "exact", "override-size", "aux-conflict"], +) +def test_concat_raises_like_dense( + parts: Callable[[LinearExpression], list[LinearExpression]], + join: JoinOptions | None, +) -> None: + require_v1() + build = "aux" if join == "outer" else "grouped" + sparse, dense = sparse_and_dense(build) + with pytest.raises(ValueError) as want: + concat(parts(dense), join) + with pytest.raises(type(want.value)), no_densify(): + concat(parts(sparse), join)