diff --git a/doc/release_notes.rst b/doc/release_notes.rst index 10b3d035..7c350633 100644 --- a/doc/release_notes.rst +++ b/doc/release_notes.rst @@ -53,6 +53,7 @@ Upcoming Version * ``@``/``dot`` against a constant matrix that holds zeros no longer densifies the result to one term per contracted member. The zero-coefficient terms are dropped, so the term dimension shrinks to the widest non-zero cell. On PyPSA's Kirchhoff Voltage Law constraint (a cycle matrix with ~3 branches per cycle) this cuts the expression from 852 to 3 terms — 284x fewer cells — which in turn shrinks the downstream ``merge``. A constant without zeros is unaffected. (`#748 `__) * ``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 `__) * Persistent snapshots of tz-aware ``DatetimeIndex`` coordinates no longer materialise an object array of ``Timestamp`` per container per capture and diff. Coordinates are stored as UTC-ns arrays with the timezone identity carried alongside, making snapshot capture ~24x and warm-start diffs ~33x faster on tz-aware models, while naive and tz-aware coordinates — and differing timezones — stay correctly unequal. (`#960 `__) **Bug fixes** diff --git a/linopy/alignment.py b/linopy/alignment.py index 88ff17f5..c566a21c 100644 --- a/linopy/alignment.py +++ b/linopy/alignment.py @@ -29,6 +29,7 @@ import numpy as np import pandas as pd import polars as pl +import scipy.sparse from numpy import arange from xarray import Coordinates, DataArray, Dataset, broadcast from xarray import align as xr_align @@ -43,6 +44,12 @@ CoordinateValidationError = ValueError # type: ignore[assignment, misc] from linopy.constants import HELPER_DIMS +from linopy.semantics import ( + _shared_dim_mismatch_message, + check_user_nan, + enforce_aux_conflict, + first_mismatched_dim, +) from linopy.types import UNLABELED_TYPES, CoordsLike, DimsLike @@ -684,6 +691,33 @@ def _matmul_operand_to_dataarray( return as_dataarray(other, coords=coords, dims=dims) +def _matmul_operand_to_matrix( + other: DataArray, + contracted: Sequence[str], + new_dims: Sequence[str], + indexes: Mapping[str, pd.Index], + aux_coords: Mapping[str, tuple[str, np.ndarray]], +) -> scipy.sparse.csr_array: + """ + Flatten a ``@`` constant to the ``(contracted, new)`` matrix of the sparse + contraction, enforcing the same rules ``*`` enforces on a constant: §5 on + NaN, §8 on the labels of the shared (contracted) dims and §11 on auxiliary + coordinates. The expression side is represented by its grid labels alone + and never broadcasts. + """ + if other.isnull().any(): + check_user_nan(op_kind="mul") + reference = Dataset(coords={d: indexes[d] for d in contracted} | dict(aux_coords)) + mismatch = first_mismatched_dim(reference, other) + if mismatch is not None: + raise ValueError(_shared_dim_mismatch_message(*mismatch)) + enforce_aux_conflict([reference, other]) + values = other.transpose(*contracted, *new_dims).values + n_contracted = int(np.prod([len(indexes[d]) for d in contracted], dtype=np.int64)) + n_new = int(np.prod([other.sizes[d] for d in new_dims], dtype=np.int64)) + return scipy.sparse.csr_array(values.reshape(n_contracted, n_new)) + + def _dims_for_positional_input( arr: Any, expected: dict[Hashable, Any], dims: DimsLike | None ) -> DimsLike | None: diff --git a/linopy/csr.py b/linopy/csr.py index df9d7d17..90d4a881 100644 --- a/linopy/csr.py +++ b/linopy/csr.py @@ -7,7 +7,10 @@ public type, different backing, akin to dask-backed xarray objects. The CSR form is canonical (duplicate variables summed, terms label-ordered) and ragged along ``_term``, so the group-size padding of issue #745 has no analog; -grouping, ``merge``/``+``/``-`` and scaling become sparse linear algebra. +grouping, ``merge``/``+``/``-``, scaling and ``@``/``dot`` (:meth:`contracted`) +become sparse linear algebra. Unlike :meth:`added`, which preserves explicit +zeros through COO, ``@``/``dot`` prunes them: cell activeness is carried by +``const`` alone (issue #925). Anything without a sparse branch expands through ``.data`` to the mathematically identical dense rectangle in canonical term layout — the reason the feature is v1-gated, where term layout is non-contractual. @@ -38,6 +41,9 @@ from linopy.expressions import LinearExpression from linopy.model import Model +CONTRACTION_CHUNK = 64 +"""Kept-axis block size of the chunked Kronecker product in ``contracted``.""" + @dataclass(frozen=True, eq=False) class Grid: @@ -396,6 +402,73 @@ def added(self, other: CSRLinearExpression) -> CSRLinearExpression: self, csr=scipy.sparse.csr_array(coo), const=const, coords=coords ) + def contracted( + self, + matrix: scipy.sparse.csr_array, + contracted_dims: Iterable[str], + new_indexes: Iterable[pd.Index], + ) -> CSRLinearExpression: + """ + Contract grid dimensions against a sparse constant (``expr @ C``). + + ``matrix`` is the constant flattened to + ``(prod(contracted shape), prod(new shape))`` in C order over + ``contracted_dims``, which are given in grid order. Each entry of + ``new_indexes`` must be named after the dim it creates, which is read + off its ``name``. The result lives on the kept grid dims followed by + ``new_indexes`` and is + ``kron(I_kept, matrix.T) @ csr``, evaluated in chunks of the kept axis + so the operator never grows with the kept size. + + The result is compact canonical form: duplicate variables summed, terms + label-ordered and explicit zeros pruned -- unlike :meth:`added`, the + sparse product drops them, so cell activeness is carried by ``const`` + alone (see issue #925). Auxiliary coordinates on kept dims propagate, + those on contracted dims drop. + """ + contracted_dims = tuple(contracted_dims) + kept = tuple(d for d in self.grid.dims if d not in contracted_dims) + target = kept + contracted_dims + source = ( + self + if self.grid.dims == target + else self.reindexed(self.grid.reordered(target)) + ) + + n_contracted = matrix.shape[0] + kept_grid = source.grid.reordered(kept) + n_kept = kept_grid.size + const = np.nan_to_num(source.const) + chunk = min(CONTRACTION_CHUNK, n_kept) + operator = scipy.sparse.kron( + scipy.sparse.eye_array(chunk), matrix.T, format="csr" + ) + + blocks = [] + const_blocks = [] + for start in range(0, n_kept, chunk): + size = min(chunk, n_kept - start) + block = ( + operator + if size == chunk + else scipy.sparse.kron( + scipy.sparse.eye_array(size), matrix.T, format="csr" + ) + ) + rows = slice(start * n_contracted, (start + size) * n_contracted) + blocks.append(block @ source.csr[rows]) + const_blocks.append(block @ const[rows]) + + indexes = kept_grid.indexes + indexes |= {str(i.name): i for i in new_indexes} + return replace( + source, + csr=scipy.sparse.csr_array(scipy.sparse.vstack(blocks, format="csr")), + const=np.concatenate(const_blocks), + grid=Grid(indexes), + coords={n: (d, v) for n, (d, v) in source.coords.items() if d in kept}, + ) + def to_dense(self) -> LinearExpression: """ Expand to the dense equivalent in canonical form: terms label-ordered, diff --git a/linopy/expressions.py b/linopy/expressions.py index e1c459eb..5385c949 100644 --- a/linopy/expressions.py +++ b/linopy/expressions.py @@ -39,7 +39,7 @@ import scipy import xarray as xr import xarray.core.groupby -from numpy import array, nan, ndarray +from numpy import array, nan from pandas.core.frame import DataFrame from pandas.core.series import Series from scipy.sparse import csc_matrix @@ -62,6 +62,7 @@ from linopy import constraints, variables from linopy.alignment import ( _matmul_operand_to_dataarray, + _matmul_operand_to_matrix, as_constant, as_dataarray, broadcast_to_coords, @@ -1523,9 +1524,13 @@ def pow(self, other: int) -> QuadraticExpression: """ return self.__pow__(other) - def dot(self, other: ndarray) -> Self | QuadraticExpression: + def dot(self, other: SideLike) -> Self | QuadraticExpression: """ Matrix multiplication with other, similar to xarray dot. + + Identical to ``@``. For a :class:`LinearExpression` that includes the + sparse contraction under v1; :class:`QuadraticExpression` always takes + the dense path. There is no per-call ``sparse=`` flag. """ return self.__matmul__(other) @@ -2539,9 +2544,20 @@ def __matmul__( ) -> LinearExpression | QuadraticExpression: """ Matrix multiplication with other, similar to xarray dot. + + Under v1, a constant ``other`` is contracted as one sparse matrix + product (:meth:`_sparse_matmul`) instead of the dense broadcast + ``(self * other).sum(dim)``. The result is then the compact canonical + form -- duplicate variables summed, terms label-ordered, explicit + zeros pruned -- so its term count may differ from the dense path's + while the values agree. A CSR-backed expression stays CSR-backed. """ other = as_constant(other) other_is_const = not isinstance(other, LinearExpression | variables.Variable) + if other_is_const and is_v1() and type(self) is LinearExpression: + sparse = self._sparse_matmul(other) + if sparse is not None: + return sparse if other_is_const: other = _matmul_operand_to_dataarray(other, self.coords, self.coord_dims) @@ -2551,6 +2567,47 @@ def __matmul__( res = res.densify_terms() return res + def _sparse_matmul(self, other: ConstantLike) -> LinearExpression | None: + """ + Contract against a constant as one sparse matrix product, skipping the + dense broadcast intermediate of ``(self * other).sum(dim)``. + + Returns None where the dense path owns the semantics: a MultiIndex on + the expression or on the operand, non-unique or MultiIndex grid + labels, a zero-size grid, an operand sharing no dimension with the + grid, and an unlabelled output dimension. The result is the compact + canonical form of :meth:`CSRLinearExpression.contracted`, so its term + count may differ from the dense path's while the values agree; it + stays CSR-backed when the input was. + """ + if is_nan_scalar(other): + check_user_nan(op_kind="mul") + if self._csr is None and _has_multiindex(self.data.indexes.values()): + return None + if not self.coord_dims: + return None + csr = self._csr or CSRLinearExpression.from_dense(self.data, self.model) + if not csr.grid.is_unique or _has_multiindex(csr.grid.indexes.values()): + return None + if csr.grid.size == 0: + return None + coords = Dataset(coords=dict(csr.grid.indexes) | dict(csr.coords)).coords + da = _matmul_operand_to_dataarray(other, coords, csr.grid.dims) + if _has_multiindex(da.indexes.values()): + return None + contracted = [d for d in csr.grid.dims if d in da.dims] + new_dims = [str(d) for d in da.dims if d not in csr.grid.dims] + if not contracted or any(d not in da.indexes for d in new_dims): + return None + matrix = _matmul_operand_to_matrix( + da, contracted, new_dims, csr.grid.indexes, csr.coords + ) + new_indexes = [da.indexes[d].rename(d) for d in new_dims] + res = csr.contracted(matrix, contracted, new_indexes) + if self._csr is not None or options["sparse_groupby"]: + return type(self)._from_csr(res, self.model) + return res.to_dense() + @property def flat(self) -> pd.DataFrame: """ @@ -3227,6 +3284,10 @@ def as_expression( return LinearExpression(obj, model) +def _has_multiindex(indexes: Iterable[pd.Index]) -> bool: + return any(isinstance(i, pd.MultiIndex) for i in indexes) + + def _aligned( csrs: list[CSRLinearExpression], join: JoinOptions | None, fill: float ) -> list[CSRLinearExpression] | None: diff --git a/linopy/semantics.py b/linopy/semantics.py index f4bc2aef..e8e584d7 100644 --- a/linopy/semantics.py +++ b/linopy/semantics.py @@ -462,7 +462,9 @@ def enforce_no_multiindex( warn_legacy(_legacy_multiindex_message(str(dim), context), stacklevel=stacklevel) -def first_mismatched_dim(a: DataArray, b: DataArray) -> tuple[str, Any, Any] | None: +def first_mismatched_dim( + a: DataArray | Dataset, b: DataArray | Dataset +) -> tuple[str, Any, Any] | None: """ Return ``(dim, a_labels, b_labels)`` for the first shared dim that disagrees on coordinate labels OR size, or ``None`` if all agree. diff --git a/test/test_csr.py b/test/test_csr.py index 8b498e56..6ea20f5a 100644 --- a/test/test_csr.py +++ b/test/test_csr.py @@ -11,19 +11,23 @@ from collections.abc import Callable from dataclasses import dataclass from pathlib import Path +from typing import Any import numpy as np import pandas as pd import polars as pl import pytest +import scipy.sparse import xarray as xr from xarray.core.types import JoinOptions import linopy from linopy import LinearExpression, Model, Variable +from linopy.constants import TERM_DIM from linopy.constraints import Constraint, ConstraintBase, CSRConstraint +from linopy.csr import CSRLinearExpression from linopy.semantics import is_v1 -from linopy.testing import assert_conequal, assert_linequal +from linopy.testing import assert_conequal, assert_linequal, assert_quadequal def require_v1() -> None: @@ -824,3 +828,394 @@ def test_cross_grid_merge_peak_memory() -> None: finally: tracemalloc.stop() assert peak < dense_rectangle_bytes / 4 + + +def cell_matrix( + expr: LinearExpression, dims: tuple[str, ...] +) -> tuple[np.ndarray, np.ndarray]: + """Per-cell coefficient row over raw variable labels, plus the flat constant.""" + ds = expr.data + nterm = ds.sizes[TERM_DIM] + vars_ = ds.vars.transpose(*dims, TERM_DIM).to_numpy().reshape(-1, nterm) + coeffs = ds.coeffs.transpose(*dims, TERM_DIM).to_numpy().reshape(-1, nterm) + active = (vars_ != -1) & ~np.isnan(coeffs) + rows = np.broadcast_to(np.arange(len(vars_))[:, None], vars_.shape) + matrix = np.zeros((len(vars_), expr.model._xCounter)) + np.add.at(matrix, (rows[active], vars_[active]), coeffs[active]) + const = np.nan_to_num(ds.const.transpose(*dims).to_numpy().reshape(-1)) + return matrix, const + + +def assert_cells_equal( + res: LinearExpression, reference: LinearExpression, dims: tuple[str, ...] +) -> None: + """Compare two expressions cell by cell, independently of the term layout.""" + got, got_const = cell_matrix(res, dims) + want, want_const = cell_matrix(reference, dims) + assert np.allclose(got, want) + assert np.allclose(got_const, want_const) + + +def assert_contracted_equal( + res: CSRLinearExpression, reference: LinearExpression, dims: tuple[str, ...] +) -> None: + """Compare a contraction with a dense reference independently of term layout.""" + assert res.grid.dims == dims + assert_cells_equal(res.to_dense(), reference, dims) + + +def flat_operand( + da: xr.DataArray, contracted_dims: tuple[str, ...], new_dims: tuple[str, ...] +) -> scipy.sparse.csr_array: + """The C-order flattening ``contracted`` expects, as a sparse matrix.""" + values = da.transpose(*contracted_dims, *new_dims).to_numpy() + n_contracted = int(np.prod([da.sizes[d] for d in contracted_dims], dtype=int)) + return scipy.sparse.csr_array(values.reshape(n_contracted, -1)) + + +def gen_csr(c: Case) -> CSRLinearExpression: + return CSRLinearExpression.from_dense((c.eff * c.gen_p).data, c.m) + + +@pytest.mark.parametrize("n_snap", [3, 70], ids=["one-chunk", "chunk-boundary"]) +@pytest.mark.parametrize("zeros", [False, True], ids=["dense-C", "sparse-C"]) +def test_contracted_partial_matches_dense(n_snap: int, zeros: bool) -> None: + require_v1() + c = base_model(n_snap=n_snap) + expr = c.eff * c.gen_p + values = np.arange(1.0, expr.data.sizes["gen"] * 2 + 1).reshape(-1, 2) + if zeros: + values[::2] = 0.0 + operand = xr.DataArray( + values, coords={"gen": expr.indexes["gen"], "loc": ["L1", "L2"]} + ) + res = gen_csr(c).contracted( + flat_operand(operand, ("gen",), ("loc",)), + ["gen"], + [pd.Index(["L1", "L2"], name="loc")], + ) + assert_contracted_equal(res, (expr * operand).sum("gen"), ("snapshot", "loc")) + + +def test_contracted_full_yields_zero_dim_grid() -> None: + require_v1() + c = base_model() + expr = c.eff * c.gen_p + operand = xr.DataArray( + np.arange( + expr.data.sizes["gen"] * expr.data.sizes["snapshot"], dtype=float + ).reshape(expr.data.sizes["gen"], -1), + coords={"gen": expr.indexes["gen"], "snapshot": expr.indexes["snapshot"]}, + ) + res = gen_csr(c).contracted( + flat_operand(operand, ("gen", "snapshot"), ()), ["gen", "snapshot"], [] + ) + assert res.grid.size == 1 + assert_contracted_equal(res, (expr * operand).sum(["gen", "snapshot"]), ()) + + +def test_contracted_with_two_new_dims() -> None: + require_v1() + c = base_model() + expr = c.eff * c.gen_p + operand = xr.DataArray( + np.arange(expr.data.sizes["gen"] * 6, dtype=float).reshape(-1, 3, 2), + coords={ + "gen": expr.indexes["gen"], + "loc": ["L1", "L2", "L3"], + "scen": [0, 1], + }, + ) + res = gen_csr(c).contracted( + flat_operand(operand, ("gen",), ("loc", "scen")), + ["gen"], + [pd.Index(["L1", "L2", "L3"], name="loc"), pd.Index([0, 1], name="scen")], + ) + assert_contracted_equal( + res, (expr * operand).sum("gen"), ("snapshot", "loc", "scen") + ) + + +def test_contracted_zero_operand_leaves_empty_rows() -> None: + require_v1() + c = base_model() + expr = c.eff * c.gen_p + operand = xr.DataArray( + np.zeros((expr.data.sizes["gen"], 2)), + coords={"gen": expr.indexes["gen"], "loc": ["L1", "L2"]}, + ) + res = gen_csr(c).contracted( + flat_operand(operand, ("gen",), ("loc",)), + ["gen"], + [pd.Index(["L1", "L2"], name="loc")], + ) + assert res.csr.nnz == 0 + assert (res.const == 0).all() + assert (res.to_dense().data.vars == -1).all() + + +def test_contracted_reorders_non_trailing_contracted_dim() -> None: + """``flow_t`` is stored as (snapshot, line); contracting ``line`` needs a reorder.""" + require_v1() + c = base_model() + expr = 1.0 * c.flow_t + csr = CSRLinearExpression.from_dense(expr.data, c.m) + assert csr.grid.dims == ("snapshot", "line") + operand = xr.DataArray( + np.arange(expr.data.sizes["line"] * 2, dtype=float).reshape(-1, 2), + coords={"line": expr.indexes["line"], "loc": ["L1", "L2"]}, + ) + res = csr.contracted( + flat_operand(operand, ("line",), ("loc",)), + ["line"], + [pd.Index(["L1", "L2"], name="loc")], + ) + assert_contracted_equal(res, (expr * operand).sum("line"), ("snapshot", "loc")) + + +LOC = pd.Index(["L1", "L2"], name="loc") + + +def loc_operand(index: pd.Index) -> xr.DataArray: + return xr.DataArray(np.ones((len(index), len(LOC))), coords=[index, LOC]) + + +def tagged_group(c: Case) -> LinearExpression: + """Generation grouped into ``group`` with ``bus``/``tag`` as auxiliary coords.""" + grouper = pd.DataFrame({"bus": c.gbus, "tag": c.gbus}) + return (1.0 * c.gen_p).groupby(grouper).sum(sparse=True, observed=True) + + +def test_contracted_keeps_aux_coords_on_kept_dims_only() -> None: + require_v1() + c = base_model() + csr = tagged_group(c)._csr + assert csr is not None + assert set(csr.coords) == {"bus", "tag"} + kept = csr.contracted( + flat_operand( + loc_operand(csr.grid.indexes["snapshot"]), ("snapshot",), ("loc",) + ), + ["snapshot"], + [LOC], + ) + assert set(kept.coords) == {"bus", "tag"} + dropped = csr.contracted( + flat_operand(loc_operand(csr.grid.indexes["group"]), ("group",), ("loc",)), + ["group"], + [LOC], + ) + assert dropped.coords == {} + + +def test_matmul_absent_cells_have_const_zero_and_no_terms() -> None: + require_v1() + c = base_model() + expr = c.eff * c.gen_p + snaps = expr.indexes["snapshot"] + expr = expr.where(xr.DataArray(snaps != 1, coords=[snaps])) + operand = loc_operand(expr.indexes["gen"]) + + res = expr @ operand + + absent = res.sel(snapshot=1) + assert (absent.data.vars == -1).all() + assert (absent.const == 0).all() + assert_cells_equal(res, (expr * operand).sum("gen"), ("snapshot", "loc")) + + +def test_matmul_prunes_explicit_zero_coefficients() -> None: + """Unlike ``added``, ``@`` drops zero terms; the cell stays active via ``const``.""" + require_v1() + c = base_model() + expr = 0.0 * c.gen_p + 1.0 + operand = loc_operand(expr.indexes["gen"]) + + res = expr @ operand + + assert (res.data.vars == -1).all() + assert (res.const == len(expr.indexes["gen"])).all() + assert_cells_equal(res, (expr * operand).sum("gen"), ("snapshot", "loc")) + + +@pytest.mark.parametrize("kind", ["nan", "mismatch", "reorder", "aux"]) +def test_matmul_operand_errors_match_multiplication(kind: str) -> None: + """§5, §8 and §11 on the constant are raised as ``*`` raises them.""" + require_v1() + c = base_model() + expr = c.eff * c.gen_p + gens = expr.indexes["gen"] + values = np.ones((len(gens), len(LOC))) + index = gens + if kind == "nan": + values[0, 0] = np.nan + elif kind == "mismatch": + index = pd.Index([f"g{i}" for i in range(len(gens))], name="gen") + elif kind == "reorder": + index = gens[::-1] + operand = xr.DataArray(values, coords=[index, LOC]) + if kind == "aux": + operand = operand.assign_coords(tag=("gen", ["X"] * len(gens))) + expr = expr.assign_coords(tag=("gen", ["Y"] * len(gens))) + + with pytest.raises(ValueError) as dense: + expr * operand + with pytest.raises(ValueError) as sparse: + expr._sparse_matmul(operand) + assert str(sparse.value) == str(dense.value) + + +@pytest.mark.parametrize("kind", ["non-unique", "multiindex"]) +def test_matmul_falls_back_to_dense_on_unsupported_labels(kind: str) -> None: + require_v1() + m = Model() + labels = ["a", "a", "b"] if kind == "non-unique" else ["a", "b", "c"] + x = m.add_variables(coords=[pd.Index(labels, name="d")], name="x") + expr = 1.0 * x + coords: Any = {"d": expr.indexes["d"]} + if kind == "multiindex": + keys = pd.MultiIndex.from_tuples( + [(0, "a"), (0, "b"), (1, "a")], names=["p", "q"] + ) + coords = xr.Coordinates.from_pandas_multiindex(keys, "d") + expr = LinearExpression(expr.data.drop_vars("d").assign_coords(coords), m) + operand = ( + xr.DataArray(np.ones((3, len(LOC))), dims=["d", "loc"]) + .assign_coords(coords) + .assign_coords(loc=LOC) + ) + + res = expr @ operand + + assert res._csr is None + assert_linequal(res, (expr * operand).sum("d")) + + +@pytest.mark.parametrize("empty", ["kept", "contracted"], ids=["kept", "contracted"]) +def test_matmul_with_zero_length_dim_matches_dense(empty: str) -> None: + """A zero-length grid dim contracts to the same empty result as the dense path.""" + require_v1() + m = Model() + gens = pd.Index(["g0", "g1"], name="gen") + empties = pd.Index([], name="empty", dtype=object) + x = m.add_variables(coords=[gens, empties], name="x") + expr = 1.0 * x + contracted = "gen" if empty == "kept" else "empty" + operand = loc_operand(expr.indexes[contracted]) + + res = expr @ operand + + reference = (expr * operand).sum(contracted) + assert res.coord_dims == ("empty" if empty == "kept" else "gen", "loc") + assert res.coord_sizes == reference.coord_sizes + assert res.data.vars.size == reference.data.vars.size == 0 + assert np.allclose(res.const, reference.const) + + +def test_matmul_dimensionless_expression_matches_dense() -> None: + """A grid-less expression has nothing to contract and stays on the dense path.""" + require_v1() + c = base_model() + expr = (1.0 * c.gen_p).sum() + operand = xr.DataArray(np.arange(2.0), coords=[LOC]) + + res = expr @ operand + + assert res._csr is None + assert_linequal(res, (expr * operand).sum([])) + + +def test_matmul_keeps_aux_coords_on_kept_dims_only() -> None: + require_v1() + c = base_model() + grouped = tagged_group(c) + assert grouped._csr is not None + indexes = grouped._csr.grid.indexes + + kept = grouped @ loc_operand(indexes["snapshot"]) + dropped = grouped @ loc_operand(indexes["group"]) + + assert {"bus", "tag"} <= set(kept.coords) + assert not {"bus", "tag"} & set(dropped.coords) + + +def test_matmul_output_dim_order_matches_dense() -> None: + require_v1() + c = base_model() + expr = c.eff * c.gen_p + operand = loc_operand(expr.indexes["gen"]) + + res = expr @ operand + + assert res.data.vars.dims == ("snapshot", "loc", TERM_DIM) + assert res.data.vars.dims == (expr * operand).sum("gen").data.vars.dims + + +def test_matmul_keeps_csr_backing_through_chain() -> None: + require_v1() + c = base_model() + snaps = c.gen_p.indexes["snapshot"] + operand = loc_operand(snaps) + gen = (c.eff * c.gen_p).groupby(c.gbus).sum(sparse=True) + flow = (1.0 * c.flow).groupby(c.bus0).sum(sparse=True) + + res = gen @ operand + assert res._csr is not None + assert_linequal(res, gen.dot(operand)) + assert_cells_equal( + res, + (c.eff * c.gen_p).groupby(c.gbus).sum() @ operand, + ("bus", "loc"), + ) + + total = res + (flow @ operand) + assert total._csr is not None + con = c.m.add_constraints(total <= 0, name="matmul", freeze=True) + assert isinstance(con, CSRConstraint) + + +@pytest.mark.legacy +def test_matmul_legacy_stays_dense() -> None: + c = base_model() + expr = c.eff * c.gen_p + operand = loc_operand(expr.indexes["gen"]) + + res = expr @ operand + + assert res._csr is None + assert_linequal(res, (expr * operand).sum("gen")) + + +def test_quadratic_matmul_stays_on_dense_path() -> None: + require_v1() + c = base_model() + quad = (1.0 * c.gen_p) * (1.0 * c.gen_p) + operand = loc_operand(quad.indexes["gen"]) + + res = quad @ operand + + assert_quadequal(res, (quad * operand).sum("gen")) + + +def test_matmul_peak_memory() -> None: + require_v1() + n_branch, n_cycle, n_snap = 1000, 300, 100 + dense_rectangle_bytes = n_snap * n_branch * n_cycle * 16 + branches = pd.Index([f"br{i}" for i in range(n_branch)], name="branch") + cycles = pd.Index(range(n_cycle), name="cycle") + snaps = pd.Index(range(n_snap), name="snapshot") + rng = np.random.default_rng(0) + incidence = xr.DataArray( + rng.uniform(size=(n_branch, n_cycle)) < 0.01, coords=[branches, cycles] + ).astype(float) + m = Model() + expr = 1.0 * m.add_variables(coords=[branches, snaps], name="flow") + + tracemalloc.start() + try: + res = expr @ incidence + _, peak = tracemalloc.get_traced_memory() + finally: + tracemalloc.stop() + assert res.coord_dims == ("snapshot", "cycle") + assert peak < dense_rectangle_bytes / 4 diff --git a/test/test_linear_expression.py b/test/test_linear_expression.py index bd2d7bff..186b7825 100644 --- a/test/test_linear_expression.py +++ b/test/test_linear_expression.py @@ -32,6 +32,7 @@ ) from linopy.constants import FACTOR_DIM, HELPER_DIMS, TERM_DIM from linopy.expressions import ScalarLinearExpression +from linopy.semantics import is_v1 from linopy.testing import assert_linequal, assert_quadequal from linopy.variables import ScalarVariable @@ -687,7 +688,41 @@ def test_matmul_expr_and_const(x: Variable, y: Variable) -> None: assert_linequal(expr.dot(const), target) -def test_matmul_contracts_only_shared_dims(z: Variable) -> None: +def backed(expr: LinearExpression, backing: str) -> LinearExpression: + """The same expression, dense or CSR-backed through an identity groupby.""" + if backing == "dense": + return expr + if not is_v1(): + pytest.skip("CSR backing requires v1 semantics") + dim = str(expr.coord_dims[-1]) + idx = expr.indexes[dim] + return expr.groupby(pd.Series(idx, index=idx, name=dim)).sum(sparse=True) + + +def assert_matmul_equal(res: LinearExpression, reference: LinearExpression) -> None: + """``assert_linequal`` up to the term layout, which ``@`` does not promise.""" + width = max(res.nterm, reference.nterm) + normed = [] + for expr in (res, reference): + pad = {TERM_DIM: (0, width - expr.nterm)} + ds = expr.data.pad(pad, constant_values={"vars": -1, "coeffs": np.nan}) + ds = ds.assign(coeffs=ds.coeffs.where(ds.vars != -1)) + normed.append(LinearExpression(ds, expr.model)) + assert_linequal(*normed) + + +def matmul_backed(expr: LinearExpression, backing: str, other: Any) -> LinearExpression: + """``@`` on the given backing, asserting the backing survives the product.""" + res = backed(expr, backing) @ other + assert (res._csr is not None) == (backing == "csr") + return res + + +backings = pytest.mark.parametrize("backing", ["dense", "csr"]) + + +@backings +def test_matmul_contracts_only_shared_dims(z: Variable, backing: str) -> None: """ A @ b contracts the genuinely shared dims and keeps the rest. @@ -703,13 +738,16 @@ def test_matmul_contracts_only_shared_dims(z: Variable) -> None: dims=["dim_1", "location"], ) - res = expr @ b + res = matmul_backed(expr, backing, b) assert set(res.coord_dims) == {"dim_0", "location"} - assert_linequal(res, (expr * b).sum("dim_1")) + assert_matmul_equal(res, (expr * b).sum("dim_1")) -def test_matmul_contracts_all_dims_when_const_covers_them(z: Variable) -> None: +@backings +def test_matmul_contracts_all_dims_when_const_covers_them( + z: Variable, backing: str +) -> None: """B covering all of a's dims (and more) contracts a's dims, keeping b's extras.""" expr = 1 * z # dims (dim_0, dim_1) b = xr.DataArray( @@ -722,13 +760,14 @@ def test_matmul_contracts_all_dims_when_const_covers_them(z: Variable) -> None: dims=["dim_0", "dim_1", "location"], ) - res = expr @ b + res = matmul_backed(expr, backing, b) assert set(res.coord_dims) == {"location"} - assert_linequal(res, (expr * b).sum(["dim_0", "dim_1"])) + assert_matmul_equal(res, (expr * b).sum(["dim_0", "dim_1"])) -def test_matmul_sparse_operand_drops_zero_terms(z: Variable) -> None: +@backings +def test_matmul_sparse_operand_drops_zero_terms(z: Variable, backing: str) -> None: """ ``@`` against a zero-containing constant compacts the term dimension to the widest non-zero cell instead of one term per contracted member (#748). @@ -740,13 +779,14 @@ def test_matmul_sparse_operand_drops_zero_terms(z: Variable) -> None: dims=["dim_1", "location"], ) - res = expr @ b + res = matmul_backed(expr, backing, b) assert res.nterm == 2 < b.sizes["dim_1"] - assert_linequal(res, (expr * b).sum("dim_1").densify_terms()) + assert_matmul_equal(res, (expr * b).sum("dim_1").densify_terms()) -def test_matmul_dense_operand_keeps_all_terms(z: Variable) -> None: +@backings +def test_matmul_dense_operand_keeps_all_terms(z: Variable, backing: str) -> None: """A constant without zeros contracts to the full term dimension, untouched.""" expr = 1 * z b = xr.DataArray( @@ -755,13 +795,14 @@ def test_matmul_dense_operand_keeps_all_terms(z: Variable) -> None: dims=["dim_1", "location"], ) - res = expr @ b + res = matmul_backed(expr, backing, b) assert res.nterm == b.sizes["dim_1"] - assert_linequal(res, (expr * b).sum("dim_1")) + assert_matmul_equal(res, (expr * b).sum("dim_1")) -def test_matmul_all_zero_operand_yields_no_terms(z: Variable) -> None: +@backings +def test_matmul_all_zero_operand_yields_no_terms(z: Variable, backing: str) -> None: """An all-zero constant contracts to the zero expression, not a crash (#748).""" expr = 1 * z b = xr.DataArray( @@ -770,23 +811,43 @@ def test_matmul_all_zero_operand_yields_no_terms(z: Variable) -> None: dims=["dim_1", "location"], ) - res = expr @ b + res = matmul_backed(expr, backing, b) - assert res.nterm == 0 assert (res.data.vars == -1).all() + assert (res.const == 0).all() -def test_matmul_full_contraction_with_zero_operand(x: Variable) -> None: +@backings +def test_matmul_full_contraction_with_zero_operand(x: Variable, backing: str) -> None: """ ``variable @ vector`` contracting away every coord dim compacts a zero-containing operand without crashing on the term-only shape (#748). """ b = xr.DataArray([2.0, 0.0], coords={"dim_0": x.indexes["dim_0"]}) - res = x @ b + res = matmul_backed(1 * x, backing, b) assert res.nterm == 1 - assert_linequal(res, (x * b).sum("dim_0").densify_terms()) + assert_matmul_equal(res, (x * b).sum("dim_0").densify_terms()) + + +@pytest.mark.v1 +@backings +def test_matmul_sums_duplicate_variables(x: Variable, backing: str) -> None: + """A variable repeated across contracted cells collapses to one summed term.""" + s = x.model.add_variables(name="dup_scalar") + expr = 1.0 * x + 2.0 * s + b = xr.DataArray( + [[1.0, 2.0], [1.0, 3.0]], + coords={"dim_0": expr.indexes["dim_0"], "loc": ["L1", "L2"]}, + ) + + res = matmul_backed(expr, backing, b) + + assert res.nterm == len(expr.indexes["dim_0"]) + 1 + assert (res.data.vars.values == [*x.labels.values, s.labels.item()]).all() + assert np.allclose(res.data.coeffs.values, [[1.0, 1.0, 4.0], [2.0, 3.0, 10.0]]) + assert (res.const == 0).all() def test_matmul_wrong_input(x: Variable, y: Variable, z: Variable) -> None: