Skip to content
Open
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
61 changes: 58 additions & 3 deletions linopy/piecewise.py
Original file line number Diff line number Diff line change
Expand Up @@ -590,9 +590,64 @@ def _breakpoints_from_slopes(

# Multi-dim case: per-entity slopes
entity_dims = [d for d in slopes_arr.dims if d != BREAKPOINT_DIM]
if len(entity_dims) != 1:
raise ValueError(
f"Expected exactly one entity dimension in slopes, got {entity_dims}"
if len(entity_dims) > 1:
if xp_arr.sizes[BREAKPOINT_DIM] == 0:
raise ValueError("Multidimensional slopes require at least one x_point")
# Align entity coordinates exactly, while breakpoint positions remain
# positional: pieces and points intentionally have different lengths.
slopes_arr, xp_arr = xr.align(
slopes_arr, xp_arr, join="exact", exclude={BREAKPOINT_DIM}
)
template = xr.broadcast(
slopes_arr.sum(BREAKPOINT_DIM),
xp_arr.sum(BREAKPOINT_DIM),
)[0]
slopes_arr = slopes_arr.broadcast_like(template)
xp_arr = xp_arr.broadcast_like(template)
for values in (slopes_arr, xp_arr):
if bool(np.isinf(values).any()):
raise ValueError("Slopes and x_points must be finite or trailing NaN")
present = values.notnull()
if bool((present & (values.isnull().cumsum(BREAKPOINT_DIM) > 0)).any()):
raise ValueError(
"Slopes and x_points may only have trailing NaN padding"
)
if bool(
(slopes_arr.count(BREAKPOINT_DIM) != xp_arr.count(BREAKPOINT_DIM) - 1).any()
):
raise ValueError("Slope count must be x_point count minus one per entity")
if bool((xp_arr.count(BREAKPOINT_DIM) < 1).any()):
raise ValueError("At least one x_point is required per entity")
if isinstance(y0, Real):
initial = xr.full_like(template, float(y0), dtype=float)
elif isinstance(y0, DataArray):
if not set(y0.dims).issubset(template.dims):
raise ValueError("y0 dimensions must be entity dimensions")
initial, _ = xr.align(y0, template, join="exact")
initial = initial.broadcast_like(template)
else:
raise TypeError("Multidimensional slopes require scalar or DataArray y0")
if not bool(np.isfinite(initial).all()):
raise ValueError("y0 must be finite")
count = xp_arr.sizes[BREAKPOINT_DIM]
xp_arr = xp_arr.assign_coords({BREAKPOINT_DIM: np.arange(count)})
slopes_arr = slopes_arr.assign_coords(
{BREAKPOINT_DIM: np.arange(slopes_arr.sizes[BREAKPOINT_DIM])}
).reindex({BREAKPOINT_DIM: np.arange(count - 1)})
widths = xp_arr.diff(BREAKPOINT_DIM).assign_coords(
{BREAKPOINT_DIM: np.arange(count - 1)}
)
cumulative = (widths * slopes_arr).cumsum(BREAKPOINT_DIM, skipna=False)
cumulative = (cumulative + initial).assign_coords(
{BREAKPOINT_DIM: np.arange(1, count)}
)
return (
xr.concat(
[initial.expand_dims({BREAKPOINT_DIM: [0]}), cumulative],
dim=BREAKPOINT_DIM,
)
.transpose(*template.dims, BREAKPOINT_DIM)
.where(xp_arr.notnull())
)
entity_dim = str(entity_dims[0])
entity_keys = slopes_arr.coords[entity_dim].values
Expand Down
141 changes: 141 additions & 0 deletions test/test_piecewise_snapshot_slopes.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,141 @@
"""
Public multidimensional slope integration contracts.

Failure list: snapshot/name coordinates are lost; shared x or y0 fails to
broadcast; leading alignment shifts a piece; ragged padding integrates a
missing piece; malformed counts, interior holes or infinity go unnoticed.
Existing tests cover one entity dimension, not snapshot/name arrays. These
independent arithmetic expectations require no production test seam.
"""

from __future__ import annotations

from typing import Literal

import numpy as np
import pandas as pd
import pytest
import xarray as xr

from linopy import Slopes
from linopy.constants import BREAKPOINT_DIM


def arrays() -> tuple[xr.DataArray, xr.DataArray]:
coords = {
"snapshot": pd.Index([2, 7]),
"name": pd.Index(["a", "b"], dtype="string"),
BREAKPOINT_DIM: [0, 1, 2],
}
x = xr.DataArray(
[[[0, 1, 3], [0, 2, 4]], [[0, 2, 3], [0, 1, np.nan]]],
dims=["snapshot", "name", BREAKPOINT_DIM],
coords=coords,
)
slopes = xr.DataArray(
[[[2, 4], [3, 5]], [[6, 8], [7, np.nan]]],
dims=x.dims,
coords={**coords, BREAKPOINT_DIM: [0, 1]},
)
return x, slopes


@pytest.mark.parametrize("leading", [False, True])
def test_snapshot_name_slopes_preserve_coordinates_and_padding(leading: bool) -> None:
x, slopes = arrays()
if leading:
first = xr.full_like(
slopes.isel({BREAKPOINT_DIM: 0}, drop=True), np.nan
).expand_dims({BREAKPOINT_DIM: [0]})
slopes = xr.concat(
[first, slopes.assign_coords({BREAKPOINT_DIM: [1, 2]})], dim=BREAKPOINT_DIM
)
y = Slopes(slopes, y0=1, align="leading" if leading else "pieces").to_breakpoints(x)
expected = xr.DataArray(
[[[1, 3, 11], [1, 7, 17]], [[1, 13, 21], [1, 8, np.nan]]],
dims=x.dims,
coords=x.coords,
)
xr.testing.assert_equal(y.transpose(*x.dims), expected)


def test_static_x_and_per_snapshot_y0_broadcast() -> None:
x, slopes = arrays()
slopes = slopes.fillna(9)
static_x = xr.DataArray([0, 1, 2], dims=[BREAKPOINT_DIM])
y0 = xr.DataArray([10, 20], dims=["snapshot"], coords={"snapshot": x.snapshot})
y = Slopes(slopes, y0=y0).to_breakpoints(static_x)
expected = xr.DataArray(
[[[10, 12, 16], [10, 13, 18]], [[20, 26, 34], [20, 27, 36]]],
dims=x.dims,
coords=x.coords,
)
xr.testing.assert_equal(y.transpose(*x.dims), expected)


@pytest.mark.parametrize(
("invalid", "message"),
[
("count", "count"),
("interior", "trailing"),
("infinite", "finite"),
("leading", "first slope"),
],
)
def test_invalid_snapshot_slopes_raise(invalid: str, message: str) -> None:
x, slopes = arrays()
align: Literal["pieces", "leading"] = "pieces"
if invalid == "count":
slopes[0, 0, 1] = np.nan
elif invalid == "interior":
x[0, 0, 1] = np.nan
elif invalid == "infinite":
slopes[0, 0, 0] = np.inf
else:
align = "leading"
with pytest.raises(ValueError, match=message):
Slopes(slopes, align=align).to_breakpoints(x)


def test_empty_slopes_with_one_x_point_keeps_initial_value() -> None:
x, slopes = arrays()
points = x.isel({BREAKPOINT_DIM: slice(0, 1)})
empty_slopes = slopes.isel({BREAKPOINT_DIM: slice(0, 0)})
result = Slopes(empty_slopes, y0=3).to_breakpoints(points)
xr.testing.assert_equal(result, xr.full_like(points, 3))


def test_empty_x_points_raise_owned_error() -> None:
x, slopes = arrays()
with pytest.raises(ValueError, match="at least one x_point"):
Slopes(slopes).to_breakpoints(x.isel({BREAKPOINT_DIM: slice(0, 0)}))


def test_mismatched_entity_coordinates_are_rejected() -> None:
x, slopes = arrays()
with pytest.raises(ValueError, match="align.*exact"):
Slopes(slopes).to_breakpoints(x.assign_coords(snapshot=[2, 8]))


@pytest.mark.parametrize(
("invalid", "message"),
[
("unknown-axis", "entity dimensions"),
("missing-coordinate", "align.*exact"),
("nan", "finite"),
("infinite", "finite"),
],
)
def test_invalid_initial_values_are_rejected(invalid: str, message: str) -> None:
x, slopes = arrays()
y0: xr.DataArray | float
if invalid == "unknown-axis":
y0 = xr.DataArray([1], dims=["unrelated"])
elif invalid == "missing-coordinate":
y0 = xr.DataArray([1], dims=["snapshot"], coords={"snapshot": [2]})
elif invalid == "nan":
y0 = np.nan
else:
y0 = np.inf
with pytest.raises(ValueError, match=message):
Slopes(slopes, y0=y0).to_breakpoints(x)
Loading