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