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
18 changes: 17 additions & 1 deletion linopy/io.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand All @@ -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``.

Expand Down Expand Up @@ -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)


Expand Down
84 changes: 84 additions & 0 deletions test/test_io.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
Loading