From 1b0441f507877ad7332d6d132a38466886385688 Mon Sep 17 00:00:00 2001 From: FBumann <117816358+FBumann@users.noreply.github.com> Date: Tue, 8 Sep 2026 10:14:41 +0200 Subject: [PATCH] feat: surface tsam's cross-column concurrency metrics AccuracyMetrics scores each clustered column on its own, so an aggregation can look accurate per column while the joint structure across them -- which values co-occur in time -- has been scrambled. That distinction matters most for exactly the multi-column, multi-slice data this wrapper exists to handle. tsam v4 added AggregationResult.concurrency; v3 has no equivalent, so _concurrency_metrics() feature-detects it and result.concurrency is None there, joining _cluster_counts() as a dual-version shim. The metrics are scalars per slice, so they concatenate along slice dims like the weighted accuracy scalars do, and they are computed on first access like accuracy. Documented in the data model rather than a notebook: the docs build runs on the locked tsam 3.4.2, where a notebook cell could only print None. Co-Authored-By: Claude Opus 5 (1M context) Claude-Session: https://claude.ai/code/session_0129FVpMdXMLjV2k19sfiWFV --- docs/data-model.md | 17 +++++ src/tsam_xarray/__init__.py | 7 +- src/tsam_xarray/_clustering.py | 2 + src/tsam_xarray/_core.py | 39 ++++++++++- src/tsam_xarray/_result.py | 66 +++++++++++++++--- test/test_concurrency.py | 123 +++++++++++++++++++++++++++++++++ 6 files changed, 243 insertions(+), 11 deletions(-) create mode 100644 test/test_concurrency.py diff --git a/docs/data-model.md b/docs/data-model.md index 110a416..4e73b46 100644 --- a/docs/data-model.md +++ b/docs/data-model.md @@ -36,6 +36,7 @@ graph LR D --- F5[".segment_durations
cluster, timestep, *slice_dims | None"] Meta --- A[".accuracy
→ AccuracyMetrics"] + Meta --- CC[".concurrency
→ ConcurrencyMetrics | None"] Meta --- C[".clustering
→ ClusteringResult"] ``` @@ -121,6 +122,22 @@ Per-column metrics as DataArrays, plus weighted scalars. | `weighted_mae` | float | Scalar MAE weighted by column weights | | `weighted_rmse_duration` | float | Scalar duration RMSE weighted by column weights | +## ConcurrencyMetrics + +Scalars measuring how well the joint structure *across* the clustered columns +— which values co-occur in time — survives aggregation, where +`AccuracyMetrics` measures each column on its own. Lower is better; both are +`NaN` when a single column is clustered. + +`result.concurrency` is `None` on tsam < 4, which does not compute these. + +| Field | Type | Description | +|-------|------|-------------| +| `correlation_error` | float | Frobenius norm of the difference between the Pearson correlation matrices of the original and the reconstructed columns | +| `rank_correlation_error` | float | The same for Spearman rank correlation, a copula proxy invariant to monotone changes in the marginals | + +With slice dims, both are DataArrays over `(*slice_dims)` rather than scalars. + ## Glossary | Term | Meaning | diff --git a/src/tsam_xarray/__init__.py b/src/tsam_xarray/__init__.py index cfbf265..e1484d8 100644 --- a/src/tsam_xarray/__init__.py +++ b/src/tsam_xarray/__init__.py @@ -3,7 +3,11 @@ from tsam_xarray._clustering import ClusteringInfo, ClusteringResult from tsam_xarray._core import aggregate from tsam_xarray._dim_names import DimNames -from tsam_xarray._result import AccuracyMetrics, AggregationResult +from tsam_xarray._result import ( + AccuracyMetrics, + AggregationResult, + ConcurrencyMetrics, +) from tsam_xarray._tuning import ( TuningResult, find_best_combination, @@ -19,6 +23,7 @@ "AggregationResult", "ClusteringInfo", "ClusteringResult", + "ConcurrencyMetrics", "DimNames", "TuningResult", "aggregate", diff --git a/src/tsam_xarray/_clustering.py b/src/tsam_xarray/_clustering.py index c6045d6..3c76520 100644 --- a/src/tsam_xarray/_clustering.py +++ b/src/tsam_xarray/_clustering.py @@ -641,6 +641,7 @@ def _apply_single( from tsam_xarray._core import ( _cluster_counts, + _concurrency_metrics, _metric_to_da, _reconstructed_to_da, _representatives_to_da, @@ -707,6 +708,7 @@ def _make_accuracy() -> AccuracyMetrics: segment_durations=seg_durations, _accuracy_factory=_make_accuracy, _reconstructed_factory=_make_reconstructed, + _concurrency_factory=lambda: _concurrency_metrics(tsam_result), original=da, clustering=clustering_info, is_transferred=True, diff --git a/src/tsam_xarray/_core.py b/src/tsam_xarray/_core.py index e791dcc..1fba35d 100644 --- a/src/tsam_xarray/_core.py +++ b/src/tsam_xarray/_core.py @@ -13,12 +13,26 @@ import xarray as xr from tsam_xarray._dim_names import DimNames -from tsam_xarray._result import AccuracyMetrics, AggregationResult +from tsam_xarray._result import AccuracyMetrics, AggregationResult, ConcurrencyMetrics Weights = dict[str, float] | dict[str, dict[str, float]] | None ClusterOn = str | Sequence[str] | dict[str, Sequence[str]] | None +def _concurrency_metrics(tsam_result: Any) -> ConcurrencyMetrics | None: + """Cross-column concurrency metrics, or None on a tsam that lacks them. + + tsam v4 added ``AggregationResult.concurrency``; v3 has no equivalent. + """ + concurrency = getattr(tsam_result, "concurrency", None) + if concurrency is None: + return None + return ConcurrencyMetrics( + correlation_error=xr.DataArray(concurrency.correlation_error), + rank_correlation_error=xr.DataArray(concurrency.rank_correlation_error), + ) + + def _cluster_counts(tsam_result: Any) -> dict[int, float]: """Per-cluster occurrence counts from a tsam result, across tsam versions. @@ -778,6 +792,7 @@ def _make_accuracy() -> AccuracyMetrics: segment_durations=seg_durations, _accuracy_factory=_make_accuracy, _reconstructed_factory=_make_reconstructed, + _concurrency_factory=lambda: _concurrency_metrics(tsam_result), original=da, clustering=clustering_info, ) @@ -819,6 +834,25 @@ def _recursive_concat(node: Any, dims: list[str]) -> xr.DataArray: return _recursive_concat(nested, slice_dims) +def _concat_concurrency( + results: list[AggregationResult], + slice_dims: list[str], + slice_coords: dict[str, Any], +) -> ConcurrencyMetrics | None: + """Concatenate per-slice concurrency metrics, or None if any slice lacks them.""" + metrics = [m for m in (r.concurrency for r in results) if m is not None] + if len(metrics) != len(results): + return None + return ConcurrencyMetrics( + correlation_error=_concat_along_dims( + [m.correlation_error for m in metrics], slice_dims, slice_coords + ), + rank_correlation_error=_concat_along_dims( + [m.rank_correlation_error for m in metrics], slice_dims, slice_coords + ), + ) + + def _concat_results( results: list[AggregationResult], slice_dims: list[str], @@ -884,6 +918,9 @@ def _acc_field(field_name: str) -> xr.DataArray: ), ), _reconstructed_factory=lambda: _field("reconstructed"), + _concurrency_factory=lambda: _concat_concurrency( + results, slice_dims, slice_coords + ), original=_field("original"), clustering=merged_clustering, is_transferred=first.is_transferred, diff --git a/src/tsam_xarray/_result.py b/src/tsam_xarray/_result.py index 452c8f8..aa763c4 100644 --- a/src/tsam_xarray/_result.py +++ b/src/tsam_xarray/_result.py @@ -17,6 +17,13 @@ from tsam_xarray._dim_names import DimNames +def _fmt_metric(da: xr.DataArray) -> str: + mean = float(da.mean()) + if da.size <= 1: + return f"{mean:.4f}" + return f"{mean:.4f} [{float(da.min()):.4f}-{float(da.max()):.4f}]" + + @dataclass(frozen=True, repr=False) class AccuracyMetrics: """Accuracy metrics from time series aggregation. @@ -45,18 +52,44 @@ class AccuracyMetrics: weighted_rmse_duration: xr.DataArray def __repr__(self) -> str: - def _fmt(da: xr.DataArray) -> str: - mean = float(da.mean()) - if da.size <= 1: - return f"{mean:.4f}" - return f"{mean:.4f} [{float(da.min()):.4f}-{float(da.max()):.4f}]" - return ( f"AccuracyMetrics(" - f"weighted_rmse={_fmt(self.weighted_rmse)}, " - f"weighted_mae={_fmt(self.weighted_mae)}, " + f"weighted_rmse={_fmt_metric(self.weighted_rmse)}, " + f"weighted_mae={_fmt_metric(self.weighted_mae)}, " f"weighted_rmse_duration=" - f"{_fmt(self.weighted_rmse_duration)})" + f"{_fmt_metric(self.weighted_rmse_duration)})" + ) + + +@dataclass(frozen=True, repr=False) +class ConcurrencyMetrics: + """Cross-column concurrency metrics from time series aggregation. + + Measures how well the joint structure across the clustered columns -- + which values co-occur in time -- survives aggregation, complementing the + per-column error in `AccuracyMetrics`. Lower is better; both values are + ``NaN`` for a single clustered column. + + Requires tsam >= 4. See `AggregationResult.concurrency`. + + Attributes: + correlation_error: Frobenius norm of the difference between the + Pearson correlation matrices of the original and the + reconstructed columns. Dims: ``(*slice_dims)`` or scalar. + rank_correlation_error: The same for the Spearman rank-correlation + matrices, a copula proxy invariant to monotone changes in the + marginals. Dims: ``(*slice_dims)`` or scalar. + """ + + correlation_error: xr.DataArray + rank_correlation_error: xr.DataArray + + def __repr__(self) -> str: + return ( + f"ConcurrencyMetrics(" + f"correlation_error={_fmt_metric(self.correlation_error)}, " + f"rank_correlation_error=" + f"{_fmt_metric(self.rank_correlation_error)})" ) @@ -81,6 +114,9 @@ class AggregationResult: Computed on first access; on a tsam that defers metric computation (v4), never reading it skips the computation entirely. + concurrency: Cross-column concurrency metrics, or + ``None`` on tsam < 4, which does not compute them. + Computed on first access, like ``accuracy``. reconstructed: Reconstructed time series (same shape and dim order as ``original``). Computed on first access, like ``accuracy``. @@ -104,6 +140,9 @@ class AggregationResult: _reconstructed_factory: Callable[[], xr.DataArray] = field( kw_only=True, repr=False, compare=False ) + _concurrency_factory: Callable[[], ConcurrencyMetrics | None] = field( + kw_only=True, repr=False, compare=False + ) @cached_property def accuracy(self) -> AccuracyMetrics: @@ -115,6 +154,15 @@ def reconstructed(self) -> xr.DataArray: """Reconstructed series on the original time axis, computed on first access.""" return self._reconstructed_factory() + @cached_property + def concurrency(self) -> ConcurrencyMetrics | None: + """Cross-column concurrency metrics, computed on first access. + + ``None`` on tsam < 4, which does not compute them. See + `ConcurrencyMetrics`. + """ + return self._concurrency_factory() + def __repr__(self) -> str: c = self.clustering slices = f", slice_dims={c.slice_dims}" if c.slice_dims else "" diff --git a/test/test_concurrency.py b/test/test_concurrency.py new file mode 100644 index 0000000..1230b91 --- /dev/null +++ b/test/test_concurrency.py @@ -0,0 +1,123 @@ +"""AggregationResult.concurrency, tsam's cross-column concurrency metrics.""" + +from __future__ import annotations + +import numpy as np +import pandas as pd +import pytest +import tsam +import xarray as xr + +import tsam_xarray +from tsam_xarray import ConcurrencyMetrics, aggregate + +HAS_CONCURRENCY = hasattr(tsam.AggregationResult, "concurrency") + +requires_concurrency = pytest.mark.skipif( + not HAS_CONCURRENCY, + reason="tsam < 4 does not compute concurrency metrics", +) + + +def _data(n_slices: int = 1, variables: list[str] | None = None) -> xr.DataArray: + if variables is None: + variables = ["a", "b", "c"] + rng = np.random.default_rng(0) + n_t = 14 * 24 + dims = ["variable", "time"] + shape: tuple[int, ...] = (len(variables), n_t) + coords: dict[str, object] = { + "variable": variables, + "time": pd.date_range("2023-01-01", periods=n_t, freq="h"), + } + if n_slices > 1: + dims = ["scenario", *dims] + shape = (n_slices, *shape) + coords["scenario"] = [f"s{i}" for i in range(n_slices)] + return xr.DataArray(rng.random(shape), dims=dims, coords=coords, name="load") + + +def _aggregate(da: xr.DataArray, n_clusters: int = 4): + return aggregate(da, time_dim="time", cluster_dim="variable", n_clusters=n_clusters) + + +def test_concurrency_metrics_is_exported(): + assert tsam_xarray.ConcurrencyMetrics is ConcurrencyMetrics + + +@pytest.mark.skipif(HAS_CONCURRENCY, reason="tsam >= 4 computes concurrency metrics") +def test_none_without_tsam_support(): + assert _aggregate(_data()).concurrency is None + + +@requires_concurrency +def test_scalar_without_slice_dims(): + concurrency = _aggregate(_data()).concurrency + + assert isinstance(concurrency, ConcurrencyMetrics) + for metric in (concurrency.correlation_error, concurrency.rank_correlation_error): + assert metric.dims == () + assert float(metric) >= 0 + + +@requires_concurrency +def test_dims_follow_slice_dims(): + da = _data(n_slices=3) + concurrency = _aggregate(da).concurrency + + for metric in (concurrency.correlation_error, concurrency.rank_correlation_error): + assert metric.dims == ("scenario",) + xr.testing.assert_identical(metric.coords["scenario"], da.coords["scenario"]) + + +@requires_concurrency +def test_exact_reconstruction_has_no_concurrency_error(): + da = _data() + n_periods = da.sizes["time"] // 24 + concurrency = _aggregate(da, n_clusters=n_periods).concurrency + + assert float(concurrency.correlation_error) == pytest.approx(0, abs=1e-9) + assert float(concurrency.rank_correlation_error) == pytest.approx(0, abs=1e-9) + + +@requires_concurrency +def test_nan_for_a_single_clustered_column(): + concurrency = _aggregate(_data(variables=["a"])).concurrency + + assert np.isnan(float(concurrency.correlation_error)) + assert np.isnan(float(concurrency.rank_correlation_error)) + + +@requires_concurrency +@pytest.mark.parametrize("n_slices", [1, 3]) +def test_deferred_until_accessed(n_slices): + result = _aggregate(_data(n_slices)) + + assert "concurrency" not in result.__dict__ + + concurrency = result.concurrency + assert "concurrency" in result.__dict__ + assert result.concurrency is concurrency + + +@requires_concurrency +@pytest.mark.parametrize("n_slices", [1, 3]) +def test_apply_reports_concurrency(n_slices): + da = _data(n_slices) + result = _aggregate(da) + transferred = result.clustering.apply(da) + + assert transferred.is_transferred + xr.testing.assert_allclose( + transferred.concurrency.correlation_error, + result.concurrency.correlation_error, + ) + + +@requires_concurrency +def test_repr_reports_both_metrics(): + text = repr(_aggregate(_data()).concurrency) + + assert text.startswith("ConcurrencyMetrics(") + assert "correlation_error=" in text + assert "rank_correlation_error=" in text