From d37e94c6c2c11ec858d09ac3d5282c39e848c543 Mon Sep 17 00:00:00 2001 From: Jonas Hoersch Date: Tue, 6 Oct 2026 10:43:40 +0200 Subject: [PATCH] feat(io): add compression option to Model.to_netcdf Applies the given encoding (e.g. {"zlib": True, "complevel": 4}) to every array in the dataset, so callers no longer need linopy's internal variable names to compress the file. Explicit `encoding` entries take precedence. Co-Authored-By: Claude Opus 5.5 --- linopy/io.py | 18 ++++++++++- test/test_io.py | 84 +++++++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 101 insertions(+), 1 deletion(-) diff --git a/linopy/io.py b/linopy/io.py index 3f2bb8f4e..205f100b2 100644 --- a/linopy/io.py +++ b/linopy/io.py @@ -1018,7 +1018,12 @@ def non_bool_dict( return {k: int(v) if isinstance(v, bool) else v for k, v in d.items()} -def to_netcdf(m: Model, *args: Any, **kwargs: Any) -> None: +DEFAULT_NETCDF_COMPRESSION: dict[str, Any] = {"zlib": True, "complevel": 3} + + +def to_netcdf( + m: Model, *args: Any, compression: dict[str, Any] | bool = True, **kwargs: Any +) -> None: """ Write out the model to a netcdf file. @@ -1028,6 +1033,11 @@ def to_netcdf(m: Model, *args: Any, **kwargs: Any) -> None: Model to write out. *args Arguments passed to ``xarray.Dataset.to_netcdf``. + compression : dict or bool, default True + Encoding applied to every array in the file. ``True`` applies + ``{"zlib": True, "complevel": 3}``, ``False`` disables + compression. It is set as a fallback in each array's + ``.encoding``. **kwargs : TYPE Keyword arguments passed to ``xarray.Dataset.to_netcdf``. @@ -1130,6 +1140,12 @@ def with_prefix(ds: xr.Dataset, prefix: str) -> xr.Dataset: for k in ds: ds[k].attrs = non_bool_dict(ds[k].attrs) + if compression is True: + compression = DEFAULT_NETCDF_COMPRESSION + if compression: + for v in ds.variables.values(): + v.encoding = compression | v.encoding + ds.to_netcdf(*args, **kwargs) diff --git a/test/test_io.py b/test/test_io.py index 67919367b..7a55916db 100644 --- a/test/test_io.py +++ b/test/test_io.py @@ -196,6 +196,90 @@ def test_model_to_netcdf_frozen_constraint(tmp_path: Path) -> None: assert_model_equal(m, p) +@pytest.mark.skipif(not HAS_NETCDF4, reason="netCDF4 not installed") +def test_model_to_netcdf_compression(tmp_path: Path) -> None: + m = Model() + x = m.add_variables(lower=0, coords=[pd.RangeIndex(2000, name="i")], name="x") + m.add_constraints(x + x.shift(i=1) >= 1, name="c", freeze=True) + m.add_constraints(x <= 10, name="d") + m.add_objective(x.sum()) + + plain, packed = tmp_path / "plain.nc", tmp_path / "packed.nc" + m.to_netcdf(plain, engine="netcdf4", compression=False) + m.to_netcdf( + packed, + engine="netcdf4", + compression={"zlib": True, "complevel": 4}, + encoding={"variables-x-lower": {"zlib": False}}, + ) + + assert packed.stat().st_size < plain.stat().st_size / 2 + with xr.open_dataset(packed, engine="netcdf4") as ds: + assert ds["variables-x-labels"].encoding["zlib"] + assert not ds["variables-x-lower"].encoding["zlib"] + assert_model_equal(m, read_netcdf(packed)) + + +@pytest.mark.skipif(not HAS_NETCDF4, reason="netCDF4 not installed") +def test_model_to_netcdf_compression_non_numeric_coords(tmp_path: Path) -> None: + m = Model() + t = pd.date_range("2030", periods=24, freq="h", name="t") + d = pd.timedelta_range("1h", periods=3, freq="h", name="d") + n = pd.Index(["north", "south"], name="n") + x = m.add_variables(lower=0, coords=[t, d, n], name="x") + m.add_constraints(x >= 1, name="c", freeze=True) + m.add_objective(x.sum()) + + fn = tmp_path / "packed.nc" + m.to_netcdf(fn, engine="netcdf4", compression={"zlib": True}) + + with xr.open_dataset(fn, engine="netcdf4") as ds: + for k in ["variables-x-t", "variables-x-d", "constraints-c-_index0"]: + assert ds[k].encoding["zlib"], k + assert_model_equal(m, read_netcdf(fn)) + + +@pytest.mark.skipif(not HAS_NETCDF4, reason="netCDF4 not installed") +def test_model_to_netcdf_default_compression(model: Model, tmp_path: Path) -> None: + def zlib(fn: Path) -> bool: + with xr.open_dataset(fn, engine="netcdf4") as ds: + return ds["variables-x-labels"].encoding["zlib"] + + model.to_netcdf(fn := tmp_path / "default.nc") + assert zlib(fn) + with xr.open_dataset(fn, engine="netcdf4") as ds: + assert ds["variables-x-labels"].encoding["complevel"] == 3 + assert_model_equal(model, read_netcdf(fn)) + + model.to_netcdf(fn := tmp_path / "off.nc", compression=False) + assert not zlib(fn) + + +@pytest.mark.skipif(not HAS_NETCDF4, reason="netCDF4 not installed") +def test_model_to_netcdf_compression_keeps_array_encoding( + model: Model, tmp_path: Path +) -> None: + user_encoding = {"dtype": "float32", "complevel": 9} + model.variables["x"].data["lower"].encoding = dict(user_encoding) + + model.to_netcdf(fn := tmp_path / "packed.nc", engine="netcdf4") + + with xr.open_dataset(fn, engine="netcdf4") as ds: + enc = ds["variables-x-lower"].encoding + assert (enc["dtype"], enc["zlib"], enc["complevel"]) == ("float32", True, 9) + assert model.variables["x"].data["lower"].encoding == user_encoding + assert model.variables["x"].data["labels"].encoding == {} + + +@pytest.mark.parametrize("compression", [True, {"zlib": True, "complevel": 4}]) +def test_model_to_netcdf_compression_scipy( + model: Model, tmp_path: Path, compression: bool | dict +) -> None: + fn = tmp_path / "scipy.nc" + model.to_netcdf(fn, engine="scipy", compression=compression) + assert_model_equal(model, read_netcdf(fn)) + + def test_model_from_netcdf_frozen_constraint_legacy_positions(tmp_path: Path) -> None: """Files written before #926 stored dense positions as CSR columns.""" from linopy.constraints import CSRConstraint