tsam-xarray 0.6.3__tar.gz → 0.6.4__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (45) hide show
  1. tsam_xarray-0.6.4/.release-please-manifest.json +3 -0
  2. {tsam_xarray-0.6.3 → tsam_xarray-0.6.4}/CHANGELOG.md +7 -0
  3. {tsam_xarray-0.6.3 → tsam_xarray-0.6.4}/PKG-INFO +1 -1
  4. {tsam_xarray-0.6.3 → tsam_xarray-0.6.4}/docs/data-model.md +28 -0
  5. {tsam_xarray-0.6.3 → tsam_xarray-0.6.4}/src/tsam_xarray/__init__.py +2 -0
  6. {tsam_xarray-0.6.3 → tsam_xarray-0.6.4}/src/tsam_xarray/_clustering.py +62 -28
  7. {tsam_xarray-0.6.3 → tsam_xarray-0.6.4}/src/tsam_xarray/_core.py +50 -24
  8. tsam_xarray-0.6.4/src/tsam_xarray/_dim_names.py +39 -0
  9. {tsam_xarray-0.6.3 → tsam_xarray-0.6.4}/src/tsam_xarray/_result.py +10 -3
  10. {tsam_xarray-0.6.3 → tsam_xarray-0.6.4}/src/tsam_xarray/_version.py +2 -2
  11. {tsam_xarray-0.6.3 → tsam_xarray-0.6.4}/test/test_aggregate.py +99 -3
  12. tsam_xarray-0.6.3/.release-please-manifest.json +0 -3
  13. {tsam_xarray-0.6.3 → tsam_xarray-0.6.4}/.github/dependabot.yml +0 -0
  14. {tsam_xarray-0.6.3 → tsam_xarray-0.6.4}/.github/workflows/ci.yaml +0 -0
  15. {tsam_xarray-0.6.3 → tsam_xarray-0.6.4}/.github/workflows/dependabot-auto-merge.yaml +0 -0
  16. {tsam_xarray-0.6.3 → tsam_xarray-0.6.4}/.github/workflows/pr-title.yaml +0 -0
  17. {tsam_xarray-0.6.3 → tsam_xarray-0.6.4}/.github/workflows/publish.yaml +0 -0
  18. {tsam_xarray-0.6.3 → tsam_xarray-0.6.4}/.github/workflows/release.yaml +0 -0
  19. {tsam_xarray-0.6.3 → tsam_xarray-0.6.4}/.gitignore +0 -0
  20. {tsam_xarray-0.6.3 → tsam_xarray-0.6.4}/.pre-commit-config.yaml +0 -0
  21. {tsam_xarray-0.6.3 → tsam_xarray-0.6.4}/.readthedocs.yaml +0 -0
  22. {tsam_xarray-0.6.3 → tsam_xarray-0.6.4}/.release-please-config.json +0 -0
  23. {tsam_xarray-0.6.3 → tsam_xarray-0.6.4}/CLAUDE.md +0 -0
  24. {tsam_xarray-0.6.3 → tsam_xarray-0.6.4}/LICENSE +0 -0
  25. {tsam_xarray-0.6.3 → tsam_xarray-0.6.4}/README.md +0 -0
  26. {tsam_xarray-0.6.3 → tsam_xarray-0.6.4}/codecov.yml +0 -0
  27. {tsam_xarray-0.6.3 → tsam_xarray-0.6.4}/docs/assets/multi-dim-input.png +0 -0
  28. {tsam_xarray-0.6.3 → tsam_xarray-0.6.4}/docs/assets/multi-dim-metrics.png +0 -0
  29. {tsam_xarray-0.6.3 → tsam_xarray-0.6.4}/docs/changelog.md +0 -0
  30. {tsam_xarray-0.6.3 → tsam_xarray-0.6.4}/docs/examples/clustering-io.ipynb +0 -0
  31. {tsam_xarray-0.6.3 → tsam_xarray-0.6.4}/docs/examples/getting-started.ipynb +0 -0
  32. {tsam_xarray-0.6.3 → tsam_xarray-0.6.4}/docs/examples/multi-dim.ipynb +0 -0
  33. {tsam_xarray-0.6.3 → tsam_xarray-0.6.4}/docs/examples/segmentation.ipynb +0 -0
  34. {tsam_xarray-0.6.3 → tsam_xarray-0.6.4}/docs/examples/tuning.ipynb +0 -0
  35. {tsam_xarray-0.6.3 → tsam_xarray-0.6.4}/docs/gen_ref_pages.py +0 -0
  36. {tsam_xarray-0.6.3 → tsam_xarray-0.6.4}/docs/index.md +0 -0
  37. {tsam_xarray-0.6.3 → tsam_xarray-0.6.4}/docs/stylesheets/extra.css +0 -0
  38. {tsam_xarray-0.6.3 → tsam_xarray-0.6.4}/mkdocs.yml +0 -0
  39. {tsam_xarray-0.6.3 → tsam_xarray-0.6.4}/pyproject.toml +0 -0
  40. {tsam_xarray-0.6.3 → tsam_xarray-0.6.4}/scripts/generate_readme_images.py +0 -0
  41. {tsam_xarray-0.6.3 → tsam_xarray-0.6.4}/src/tsam_xarray/_sample_data.py +0 -0
  42. {tsam_xarray-0.6.3 → tsam_xarray-0.6.4}/src/tsam_xarray/_tuning.py +0 -0
  43. {tsam_xarray-0.6.3 → tsam_xarray-0.6.4}/test/conftest.py +0 -0
  44. {tsam_xarray-0.6.3 → tsam_xarray-0.6.4}/test/test_parametrized.py +0 -0
  45. {tsam_xarray-0.6.3 → tsam_xarray-0.6.4}/test/test_tuning.py +0 -0
@@ -0,0 +1,3 @@
1
+ {
2
+ ".": "0.6.4"
3
+ }
@@ -1,5 +1,12 @@
1
1
  # Changelog
2
2
 
3
+ ## [0.6.4](https://github.com/FBumann/tsam_xarray/compare/v0.6.3...v0.6.4) (2026-07-22)
4
+
5
+
6
+ ### Features
7
+
8
+ * make output dimension names configurable via DimNames ([#101](https://github.com/FBumann/tsam_xarray/issues/101)) ([558fb9d](https://github.com/FBumann/tsam_xarray/commit/558fb9da4014d4bdd05515f79fef5c542c17239e)), closes [#99](https://github.com/FBumann/tsam_xarray/issues/99)
9
+
3
10
  ## [0.6.3](https://github.com/FBumann/tsam_xarray/compare/v0.6.2...v0.6.3) (2026-07-22)
4
11
 
5
12
 
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: tsam_xarray
3
- Version: 0.6.3
3
+ Version: 0.6.4
4
4
  Summary: Lightweight xarray wrapper for tsam time series aggregation
5
5
  License-Expression: MIT
6
6
  License-File: LICENSE
@@ -41,6 +41,34 @@ graph LR
41
41
 
42
42
  </div>
43
43
 
44
+ ### Custom output dimension names
45
+
46
+ `aggregate()` adds four structural dimensions to its results that do not
47
+ exist in the input: `cluster`, `timestep`, `period` (in
48
+ `cluster_assignments`), and `segment` (segmented runs). These names are
49
+ reserved — an input dimension of the same name raises an error.
50
+
51
+ Pass a `DimNames` to rename them, e.g. when a caller already has a
52
+ `period` dimension (multi-period optimization models):
53
+
54
+ ```python
55
+ from tsam_xarray import DimNames, aggregate
56
+
57
+ result = aggregate(
58
+ da, # has a slice dim literally named "period"
59
+ time_dim="time",
60
+ cluster_dim="variable",
61
+ n_clusters=8,
62
+ dim_names=DimNames(period="original_period"),
63
+ )
64
+ result.cluster_assignments.dims # ("original_period", ...)
65
+ ```
66
+
67
+ The resolved names are stored on `ClusteringResult`, so `apply()`,
68
+ `disaggregate()`, and the JSON round-trip all reproduce them. `dim_names`
69
+ defaults to `None`, which keeps today's names. The chosen names must be
70
+ unique and must not collide with any input dimension.
71
+
44
72
  ## ClusteringResult
45
73
 
46
74
  The reusable part — knows *how* the time series was clustered,
@@ -2,6 +2,7 @@
2
2
 
3
3
  from tsam_xarray._clustering import ClusteringInfo, ClusteringResult
4
4
  from tsam_xarray._core import aggregate
5
+ from tsam_xarray._dim_names import DimNames
5
6
  from tsam_xarray._result import AccuracyMetrics, AggregationResult
6
7
  from tsam_xarray._tuning import (
7
8
  TuningResult,
@@ -18,6 +19,7 @@ __all__ = [
18
19
  "AggregationResult",
19
20
  "ClusteringInfo",
20
21
  "ClusteringResult",
22
+ "DimNames",
21
23
  "TuningResult",
22
24
  "aggregate",
23
25
  "find_best_combination",
@@ -19,6 +19,7 @@ from tsam_xarray._core import (
19
19
  _resolve_cluster_dim,
20
20
  _segment_durations_to_da,
21
21
  )
22
+ from tsam_xarray._dim_names import DimNames
22
23
 
23
24
 
24
25
  @dataclass(frozen=True, repr=False)
@@ -53,12 +54,15 @@ class ClusteringResult:
53
54
  segment_centers: Representative timestep per segment,
54
55
  or ``None``.
55
56
  Dims: ``(cluster, segment, *slice_dims)``.
57
+ dim_names: Names of the structural output dimensions.
58
+ See `DimNames`.
56
59
  """
57
60
 
58
61
  time_dim: str
59
62
  cluster_dim: list[str]
60
63
  slice_dims: list[str]
61
64
  clusterings: dict[tuple[Hashable, ...], tsam.ClusteringResult]
65
+ dim_names: DimNames = field(default_factory=DimNames)
62
66
  _cache: dict[str, Any] = field(
63
67
  default_factory=dict, repr=False, init=False, compare=False
64
68
  )
@@ -125,14 +129,19 @@ class ClusteringResult:
125
129
  def _build_assignments(self) -> xr.DataArray:
126
130
  if not self.slice_dims:
127
131
  cr = self.clusterings[()]
128
- return xr.DataArray(list(cr.cluster_assignments), dims=["period"])
132
+ return xr.DataArray(
133
+ list(cr.cluster_assignments), dims=[self.dim_names.period]
134
+ )
129
135
 
130
136
  import itertools
131
137
 
132
138
  sc = self._slice_coords
133
139
  keys = list(itertools.product(*(sc[d] for d in self.slice_dims)))
134
140
  arrays = [
135
- xr.DataArray(list(self.clusterings[k].cluster_assignments), dims=["period"])
141
+ xr.DataArray(
142
+ list(self.clusterings[k].cluster_assignments),
143
+ dims=[self.dim_names.period],
144
+ )
136
145
  for k in keys
137
146
  ]
138
147
  return _concat_along_dims(arrays, self.slice_dims, sc)
@@ -153,8 +162,8 @@ class ClusteringResult:
153
162
  counts = np.bincount(cr.cluster_assignments, minlength=cr.n_clusters)
154
163
  return xr.DataArray(
155
164
  counts,
156
- dims=["cluster"],
157
- coords={"cluster": np.arange(cr.n_clusters)},
165
+ dims=[self.dim_names.cluster],
166
+ coords={self.dim_names.cluster: np.arange(cr.n_clusters)},
158
167
  )
159
168
 
160
169
  if not self.slice_dims:
@@ -180,18 +189,24 @@ class ClusteringResult:
180
189
 
181
190
  def _build_segment_durations(self) -> xr.DataArray | None:
182
191
  if not self.slice_dims:
183
- return _segment_durations_to_da(self.clusterings[()].segment_durations)
192
+ return _segment_durations_to_da(
193
+ self.clusterings[()].segment_durations, self.dim_names
194
+ )
184
195
 
185
196
  import itertools
186
197
 
187
198
  sc = self._slice_coords
188
199
  keys = list(itertools.product(*(sc[d] for d in self.slice_dims)))
189
- first = _segment_durations_to_da(self.clusterings[keys[0]].segment_durations)
200
+ first = _segment_durations_to_da(
201
+ self.clusterings[keys[0]].segment_durations, self.dim_names
202
+ )
190
203
  if first is None:
191
204
  return None
192
205
  das: list[xr.DataArray] = [first]
193
206
  for k in keys[1:]:
194
- da = _segment_durations_to_da(self.clusterings[k].segment_durations)
207
+ da = _segment_durations_to_da(
208
+ self.clusterings[k].segment_durations, self.dim_names
209
+ )
195
210
  if da is None:
196
211
  msg = (
197
212
  f"Slice {k} has no segment durations but the first "
@@ -220,8 +235,8 @@ class ClusteringResult:
220
235
  raise ValueError(msg)
221
236
  return xr.DataArray(
222
237
  list(centers),
223
- dims=["cluster"],
224
- coords={"cluster": np.arange(cr.n_clusters)},
238
+ dims=[self.dim_names.cluster],
239
+ coords={self.dim_names.cluster: np.arange(cr.n_clusters)},
225
240
  )
226
241
 
227
242
  if not self.slice_dims:
@@ -251,10 +266,10 @@ class ClusteringResult:
251
266
  return None
252
267
  return xr.DataArray(
253
268
  np.array(cr.segment_assignments),
254
- dims=["cluster", "timestep"],
269
+ dims=[self.dim_names.cluster, self.dim_names.timestep],
255
270
  coords={
256
- "cluster": np.arange(cr.n_clusters),
257
- "timestep": np.arange(cr.n_timesteps_per_period),
271
+ self.dim_names.cluster: np.arange(cr.n_clusters),
272
+ self.dim_names.timestep: np.arange(cr.n_timesteps_per_period),
258
273
  },
259
274
  )
260
275
 
@@ -298,10 +313,10 @@ class ClusteringResult:
298
313
  n_segments = cr.n_segments or len(cr.segment_centers[0])
299
314
  return xr.DataArray(
300
315
  np.array(cr.segment_centers),
301
- dims=["cluster", "segment"],
316
+ dims=[self.dim_names.cluster, self.dim_names.segment],
302
317
  coords={
303
- "cluster": np.arange(cr.n_clusters),
304
- "segment": np.arange(n_segments),
318
+ self.dim_names.cluster: np.arange(cr.n_clusters),
319
+ self.dim_names.segment: np.arange(n_segments),
305
320
  },
306
321
  )
307
322
 
@@ -367,7 +382,7 @@ class ClusteringResult:
367
382
 
368
383
  if not slice_dims:
369
384
  cr = self.clusterings[()]
370
- return _apply_single(da, cr, td, cd, tsam_kwargs)
385
+ return _apply_single(da, cr, td, cd, tsam_kwargs, self.dim_names)
371
386
 
372
387
  import itertools
373
388
 
@@ -380,7 +395,7 @@ class ClusteringResult:
380
395
  sel = dict(zip(slice_dims, key, strict=True))
381
396
  da_slice = da.sel(sel)
382
397
  cr = _lookup_clustering(self.clusterings, key)
383
- r = _apply_single(da_slice, cr, td, cd, tsam_kwargs)
398
+ r = _apply_single(da_slice, cr, td, cd, tsam_kwargs, self.dim_names)
384
399
  results.append(r)
385
400
 
386
401
  return _concat_results(results, slice_dims, slice_coords, slice_keys)
@@ -408,7 +423,7 @@ class ClusteringResult:
408
423
  """
409
424
  slice_dims = self.slice_dims
410
425
  if not slice_dims:
411
- return _disaggregate_single(self.clusterings[()], data)
426
+ return _disaggregate_single(self.clusterings[()], data, self.dim_names)
412
427
 
413
428
  import itertools
414
429
 
@@ -419,7 +434,7 @@ class ClusteringResult:
419
434
  sel = dict(zip(slice_dims, key, strict=True))
420
435
  data_slice = data.sel(sel)
421
436
  cr = _lookup_clustering(self.clusterings, key)
422
- results.append(_disaggregate_single(cr, data_slice))
437
+ results.append(_disaggregate_single(cr, data_slice, self.dim_names))
423
438
 
424
439
  return _concat_along_dims(results, slice_dims, slice_coords)
425
440
 
@@ -442,6 +457,12 @@ class ClusteringResult:
442
457
  "time_dim": self.time_dim,
443
458
  "cluster_dim": self.cluster_dim,
444
459
  "slice_dims": self.slice_dims,
460
+ "dim_names": {
461
+ "cluster": self.dim_names.cluster,
462
+ "timestep": self.dim_names.timestep,
463
+ "period": self.dim_names.period,
464
+ "segment": self.dim_names.segment,
465
+ },
445
466
  "clusterings": entries,
446
467
  }
447
468
 
@@ -486,11 +507,15 @@ class ClusteringResult:
486
507
  key = tuple(entry["key"])
487
508
  clusterings[key] = tsam.ClusteringResult.from_dict(entry["clustering"])
488
509
 
510
+ dim_names_data = data.get("dim_names")
511
+ dim_names = DimNames(**dim_names_data) if dim_names_data else DimNames()
512
+
489
513
  return cls(
490
514
  time_dim=data["time_dim"],
491
515
  cluster_dim=data["cluster_dim"],
492
516
  slice_dims=data.get("slice_dims", []),
493
517
  clusterings=clusterings,
518
+ dim_names=dim_names,
494
519
  )
495
520
 
496
521
  @classmethod
@@ -603,6 +628,7 @@ def _apply_single(
603
628
  time_dim: str,
604
629
  col_dims: list[str],
605
630
  tsam_kwargs: dict[str, Any],
631
+ dim_names: DimNames,
606
632
  ) -> Any:
607
633
  """Apply a single ClusteringResult to a DataArray."""
608
634
  import pandas as pd
@@ -620,18 +646,22 @@ def _apply_single(
620
646
  cr = _drop_missing_weights(cr, df)
621
647
  tsam_result = cr.apply(df, **tsam_kwargs)
622
648
 
623
- typical = _representatives_to_da(tsam_result.cluster_representatives, col_dims)
649
+ typical = _representatives_to_da(
650
+ tsam_result.cluster_representatives, col_dims, dim_names
651
+ )
624
652
  reconstructed = _reconstructed_to_da(tsam_result.reconstructed, time_dim, col_dims)
625
653
 
626
654
  cw = tsam_result.cluster_weights
627
655
  cluster_ids = np.array(sorted(cw.keys()))
628
656
  cluster_weights_da = xr.DataArray(
629
657
  np.array([cw[k] for k in cluster_ids]),
630
- dims=["cluster"],
631
- coords={"cluster": cluster_ids},
658
+ dims=[dim_names.cluster],
659
+ coords={dim_names.cluster: cluster_ids},
632
660
  )
633
661
 
634
- assignments_da = xr.DataArray(tsam_result.cluster_assignments, dims=["period"])
662
+ assignments_da = xr.DataArray(
663
+ tsam_result.cluster_assignments, dims=[dim_names.period]
664
+ )
635
665
 
636
666
  col_names: list[str] | None = None
637
667
  if isinstance(df.columns, pd.MultiIndex):
@@ -650,13 +680,14 @@ def _apply_single(
650
680
  ),
651
681
  )
652
682
 
653
- seg_durations = _segment_durations_to_da(tsam_result.segment_durations)
683
+ seg_durations = _segment_durations_to_da(tsam_result.segment_durations, dim_names)
654
684
 
655
685
  clustering_info = ClusteringResult(
656
686
  time_dim=time_dim,
657
687
  cluster_dim=col_dims,
658
688
  slice_dims=[],
659
689
  clusterings={(): tsam_result.clustering},
690
+ dim_names=dim_names,
660
691
  )
661
692
 
662
693
  return AggregationResult(
@@ -675,18 +706,21 @@ def _apply_single(
675
706
  def _disaggregate_single(
676
707
  cr: tsam.ClusteringResult,
677
708
  data: xr.DataArray,
709
+ dim_names: DimNames,
678
710
  ) -> xr.DataArray:
679
711
  """Disaggregate a single (non-sliced) DataArray using a ClusteringResult.
680
712
 
681
713
  Relies on tsam's ``cr.disaggregate()`` to return a DataFrame indexed
682
714
  by the original ``DatetimeIndex`` stored on the clustering.
683
715
  """
684
- other_dims = [str(d) for d in data.dims if d not in ("cluster", "timestep")]
685
- ordered = data.transpose("cluster", "timestep", *other_dims)
716
+ cluster_dim = dim_names.cluster
717
+ timestep_dim = dim_names.timestep
718
+ other_dims = [str(d) for d in data.dims if d not in (cluster_dim, timestep_dim)]
719
+ ordered = data.transpose(cluster_dim, timestep_dim, *other_dims)
686
720
 
687
- clusters = ordered.coords["cluster"].values
721
+ clusters = ordered.coords[cluster_dim].values
688
722
  n_clusters = len(clusters)
689
- n_timesteps = ordered.sizes["timestep"]
723
+ n_timesteps = ordered.sizes[timestep_dim]
690
724
  other_sizes = ordered.shape[2:]
691
725
 
692
726
  flat = ordered.values.reshape(n_clusters * n_timesteps, -1)
@@ -12,6 +12,7 @@ import pandas as pd
12
12
  import tsam
13
13
  import xarray as xr
14
14
 
15
+ from tsam_xarray._dim_names import DimNames
15
16
  from tsam_xarray._result import AccuracyMetrics, AggregationResult
16
17
 
17
18
  Weights = dict[str, float] | dict[str, dict[str, float]] | None
@@ -26,6 +27,7 @@ def aggregate(
26
27
  n_clusters: int,
27
28
  weights: Weights = None,
28
29
  cluster_on: ClusterOn = None,
30
+ dim_names: DimNames | None = None,
29
31
  **tsam_kwargs: Any,
30
32
  ) -> AggregationResult:
31
33
  """Aggregate an xarray DataArray using tsam.
@@ -89,13 +91,21 @@ def aggregate(
89
91
  extreme periods are identified only on the
90
92
  clustered-on columns.
91
93
 
94
+ dim_names: Names for the structural output dimensions
95
+ (``cluster``, ``timestep``, ``period``, ``segment``).
96
+ ``None`` (default) keeps today's names. Override to
97
+ avoid collisions with the caller's own dimension
98
+ names. See `DimNames`.
99
+
92
100
  **tsam_kwargs: Additional keyword arguments passed to
93
101
  ``tsam.aggregate()``.
94
102
  """
103
+ resolved_dim_names = dim_names if dim_names is not None else DimNames()
95
104
  _validate_time_dim(da, time_dim)
96
105
  col_dims = _resolve_cluster_dim(cluster_dim)
97
106
  slice_dims = _infer_slice_dims(da, time_dim, col_dims)
98
107
  _validate(da, time_dim, col_dims, slice_dims)
108
+ _validate_dim_names(resolved_dim_names, time_dim, col_dims, slice_dims)
99
109
  da = _validate_data(da, time_dim, col_dims, slice_dims)
100
110
  _validate_no_cluster_config_weights(tsam_kwargs)
101
111
  per_dim_weights = _normalize_weights(weights, da, col_dims)
@@ -111,6 +121,7 @@ def aggregate(
111
121
  per_dim_weights,
112
122
  active_coords,
113
123
  tsam_kwargs,
124
+ resolved_dim_names,
114
125
  )
115
126
 
116
127
  slice_coords = {d: da.coords[d].values for d in slice_dims}
@@ -129,6 +140,7 @@ def aggregate(
129
140
  per_dim_weights,
130
141
  active_coords,
131
142
  tsam_kwargs,
143
+ resolved_dim_names,
132
144
  )
133
145
  results.append(r)
134
146
 
@@ -255,7 +267,22 @@ def _validate(
255
267
  raise ValueError(msg)
256
268
 
257
269
 
258
- _RESERVED_DIMS = {"cluster", "timestep", "period"}
270
+ def _validate_dim_names(
271
+ dim_names: DimNames,
272
+ time_dim: str,
273
+ col_dims: list[str],
274
+ slice_dims: list[str],
275
+ ) -> None:
276
+ """Reject output dimension names that collide with input dimensions."""
277
+ input_dims = {time_dim, *col_dims, *slice_dims}
278
+ conflict = input_dims & set(dim_names.as_tuple())
279
+ if conflict:
280
+ msg = (
281
+ f"Output dimension names {conflict} collide with input dimension "
282
+ "names. Rename them in your input DataArray, or pass a DimNames "
283
+ "with different names via dim_names=."
284
+ )
285
+ raise ValueError(msg)
259
286
 
260
287
 
261
288
  def _validate_data(
@@ -268,17 +295,6 @@ def _validate_data(
268
295
 
269
296
  Returns the (possibly computed) DataArray.
270
297
  """
271
- # Reserved dimension names
272
- all_user_dims = {time_dim, *col_dims, *slice_dims}
273
- reserved_conflict = all_user_dims & _RESERVED_DIMS
274
- if reserved_conflict:
275
- msg = (
276
- f"Dimension names {reserved_conflict} are reserved by "
277
- "tsam_xarray for output dimensions. Rename them in "
278
- "your input DataArray."
279
- )
280
- raise ValueError(msg)
281
-
282
298
  # Dask arrays — compute before other checks
283
299
  if hasattr(da.data, "dask"):
284
300
  warnings.warn(
@@ -368,6 +384,7 @@ def _to_dataframe(
368
384
  def _representatives_to_da(
369
385
  df: pd.DataFrame,
370
386
  col_dims: list[str],
387
+ dim_names: DimNames,
371
388
  ) -> xr.DataArray:
372
389
  """Convert cluster_representatives DataFrame to DataArray."""
373
390
  df = df.copy()
@@ -375,7 +392,7 @@ def _representatives_to_da(
375
392
  # Without: 2 levels: (cluster, timestep)
376
393
  if isinstance(df.index, pd.MultiIndex) and df.index.nlevels == 3:
377
394
  df.index = df.index.droplevel(2) # drop segment_duration
378
- df.index.names = ["cluster", "timestep"]
395
+ df.index.names = [dim_names.cluster, dim_names.timestep]
379
396
 
380
397
  if not col_dims:
381
398
  clusters = df.index.get_level_values(0).unique()
@@ -383,8 +400,8 @@ def _representatives_to_da(
383
400
  values = df.values.squeeze(axis=1).reshape(len(clusters), len(timesteps))
384
401
  return xr.DataArray(
385
402
  values,
386
- dims=["cluster", "timestep"],
387
- coords={"cluster": clusters, "timestep": timesteps},
403
+ dims=[dim_names.cluster, dim_names.timestep],
404
+ coords={dim_names.cluster: clusters, dim_names.timestep: timesteps},
388
405
  )
389
406
 
390
407
  stacked = df.stack(df.columns.names, future_stack=True)
@@ -394,6 +411,7 @@ def _representatives_to_da(
394
411
 
395
412
  def _segment_durations_to_da(
396
413
  raw_durations: tuple[tuple[int, ...], ...] | None,
414
+ dim_names: DimNames,
397
415
  ) -> xr.DataArray | None:
398
416
  """Convert tsam segment_durations to DataArray."""
399
417
  if raw_durations is None:
@@ -401,10 +419,10 @@ def _segment_durations_to_da(
401
419
  data = np.array(raw_durations) # (n_clusters, n_segments)
402
420
  return xr.DataArray(
403
421
  data,
404
- dims=["cluster", "timestep"],
422
+ dims=[dim_names.cluster, dim_names.timestep],
405
423
  coords={
406
- "cluster": np.arange(data.shape[0]),
407
- "timestep": np.arange(data.shape[1]),
424
+ dim_names.cluster: np.arange(data.shape[0]),
425
+ dim_names.timestep: np.arange(data.shape[1]),
408
426
  },
409
427
  )
410
428
 
@@ -624,6 +642,7 @@ def _aggregate_single(
624
642
  weights: dict[str, dict[str, float]] | None,
625
643
  active_coords: dict[str, set[str]] | None,
626
644
  tsam_kwargs: dict[str, Any],
645
+ dim_names: DimNames,
627
646
  ) -> AggregationResult:
628
647
  """Run a single tsam aggregation on a DataArray.
629
648
 
@@ -667,7 +686,7 @@ def _aggregate_single(
667
686
  apply_kwargs = {k: v for k, v in tsam_kwargs.items() if k in _APPLY_KWARGS}
668
687
  tsam_result = clustering.apply(df, **apply_kwargs)
669
688
 
670
- return _result_from_tsam(tsam_result, da, df, time_dim, col_dims)
689
+ return _result_from_tsam(tsam_result, da, df, time_dim, col_dims, dim_names)
671
690
 
672
691
 
673
692
  def _result_from_tsam(
@@ -676,9 +695,12 @@ def _result_from_tsam(
676
695
  df: pd.DataFrame,
677
696
  time_dim: str,
678
697
  col_dims: list[str],
698
+ dim_names: DimNames,
679
699
  ) -> AggregationResult:
680
700
  """Build an AggregationResult from a tsam aggregation result."""
681
- typical = _representatives_to_da(tsam_result.cluster_representatives, col_dims)
701
+ typical = _representatives_to_da(
702
+ tsam_result.cluster_representatives, col_dims, dim_names
703
+ )
682
704
  reconstructed = _reconstructed_to_da(tsam_result.reconstructed, time_dim, col_dims)
683
705
  reconstructed = reconstructed.transpose(*da.dims).reindex_like(da)
684
706
 
@@ -686,11 +708,13 @@ def _result_from_tsam(
686
708
  cluster_ids = np.array(sorted(cw.keys()))
687
709
  cluster_weights_da = xr.DataArray(
688
710
  np.array([cw[k] for k in cluster_ids]),
689
- dims=["cluster"],
690
- coords={"cluster": cluster_ids},
711
+ dims=[dim_names.cluster],
712
+ coords={dim_names.cluster: cluster_ids},
691
713
  )
692
714
 
693
- assignments_da = xr.DataArray(tsam_result.cluster_assignments, dims=["period"])
715
+ assignments_da = xr.DataArray(
716
+ tsam_result.cluster_assignments, dims=[dim_names.period]
717
+ )
694
718
 
695
719
  col_names: list[str] | None = None
696
720
  if isinstance(df.columns, pd.MultiIndex):
@@ -709,7 +733,7 @@ def _result_from_tsam(
709
733
  ),
710
734
  )
711
735
 
712
- seg_durations = _segment_durations_to_da(tsam_result.segment_durations)
736
+ seg_durations = _segment_durations_to_da(tsam_result.segment_durations, dim_names)
713
737
 
714
738
  from tsam_xarray._clustering import ClusteringResult
715
739
 
@@ -718,6 +742,7 @@ def _result_from_tsam(
718
742
  cluster_dim=col_dims,
719
743
  slice_dims=[],
720
744
  clusterings={(): tsam_result.clustering},
745
+ dim_names=dim_names,
721
746
  )
722
747
 
723
748
  return AggregationResult(
@@ -804,6 +829,7 @@ def _concat_results(
804
829
  cluster_dim=first.clustering.cluster_dim,
805
830
  slice_dims=slice_dims,
806
831
  clusterings=merged_clusterings,
832
+ dim_names=first.clustering.dim_names,
807
833
  )
808
834
 
809
835
  return AggregationResult(
@@ -0,0 +1,39 @@
1
+ """Configurable output dimension names for tsam_xarray."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from dataclasses import dataclass
6
+
7
+
8
+ @dataclass(frozen=True)
9
+ class DimNames:
10
+ """Names of the structural output dimensions produced by aggregation.
11
+
12
+ tsam_xarray adds four dimensions to its results that do not exist in the
13
+ input: the cluster/representative axis, the intra-period timestep axis,
14
+ the original-period axis (in ``cluster_assignments``), and the segment
15
+ axis (segmented runs). By default these are ``cluster``, ``timestep``,
16
+ ``period``, and ``segment``; override them when they would collide with
17
+ the caller's own dimension names.
18
+
19
+ Attributes:
20
+ cluster: Cluster/representative axis.
21
+ timestep: Intra-period timestep axis.
22
+ period: Original-period axis (in ``cluster_assignments``).
23
+ segment: Segment axis (segmented runs).
24
+ """
25
+
26
+ cluster: str = "cluster"
27
+ timestep: str = "timestep"
28
+ period: str = "period"
29
+ segment: str = "segment"
30
+
31
+ def __post_init__(self) -> None:
32
+ names = self.as_tuple()
33
+ if len(set(names)) != len(names):
34
+ msg = f"DimNames must be unique, got {names}"
35
+ raise ValueError(msg)
36
+
37
+ def as_tuple(self) -> tuple[str, str, str, str]:
38
+ """The four names as a tuple, in declaration order."""
39
+ return (self.cluster, self.timestep, self.period, self.segment)
@@ -10,6 +10,7 @@ import xarray as xr
10
10
 
11
11
  if TYPE_CHECKING:
12
12
  from tsam_xarray._clustering import ClusteringResult
13
+ from tsam_xarray._dim_names import DimNames
13
14
 
14
15
 
15
16
  @dataclass(frozen=True, repr=False)
@@ -103,15 +104,20 @@ class AggregationResult:
103
104
  f"weighted_rmse={float(self.accuracy.weighted_rmse.mean()):.4f})"
104
105
  )
105
106
 
107
+ @property
108
+ def dim_names(self) -> DimNames:
109
+ """Names of the structural output dimensions. See `DimNames`."""
110
+ return self.clustering.dim_names
111
+
106
112
  @property
107
113
  def n_clusters(self) -> int:
108
114
  """Number of cluster representative clusters."""
109
- return int(self.cluster_weights.sizes["cluster"])
115
+ return int(self.cluster_weights.sizes[self.dim_names.cluster])
110
116
 
111
117
  @property
112
118
  def n_timesteps_per_period(self) -> int:
113
119
  """Number of timesteps per cluster representative."""
114
- return int(self.cluster_representatives.sizes["timestep"])
120
+ return int(self.cluster_representatives.sizes[self.dim_names.timestep])
115
121
 
116
122
  @property
117
123
  def n_segments(self) -> int | None:
@@ -264,6 +270,7 @@ class AggregationResult:
264
270
  cluster_dim=self.clustering.cluster_dim,
265
271
  slice_dims=[],
266
272
  clusterings={(): cr},
273
+ dim_names=self.clustering.dim_names,
267
274
  ),
268
275
  )
269
276
 
@@ -272,4 +279,4 @@ class AggregationResult:
272
279
  from tsam_xarray._clustering import _disaggregate_single
273
280
 
274
281
  cr = self.clustering.clusterings[()]
275
- return _disaggregate_single(cr, data)
282
+ return _disaggregate_single(cr, data, self.clustering.dim_names)
@@ -18,7 +18,7 @@ version_tuple: tuple[int | str, ...]
18
18
  commit_id: str | None
19
19
  __commit_id__: str | None
20
20
 
21
- __version__ = version = '0.6.3'
22
- __version_tuple__ = version_tuple = (0, 6, 3)
21
+ __version__ = version = '0.6.4'
22
+ __version_tuple__ = version_tuple = (0, 6, 4)
23
23
 
24
24
  __commit_id__ = commit_id = None
@@ -833,7 +833,7 @@ class TestDataValidation:
833
833
  dims=["time", "cluster"],
834
834
  coords={"time": time, "cluster": ["a", "b"]},
835
835
  )
836
- with pytest.raises(ValueError, match="reserved"):
836
+ with pytest.raises(ValueError, match="collide"):
837
837
  tsam_xarray.aggregate(
838
838
  da, time_dim="time", cluster_dim="cluster", n_clusters=4
839
839
  )
@@ -845,7 +845,7 @@ class TestDataValidation:
845
845
  dims=["time", "timestep"],
846
846
  coords={"time": time, "timestep": [0, 1]},
847
847
  )
848
- with pytest.raises(ValueError, match="reserved"):
848
+ with pytest.raises(ValueError, match="collide"):
849
849
  tsam_xarray.aggregate(
850
850
  da, time_dim="time", cluster_dim="timestep", n_clusters=4
851
851
  )
@@ -861,7 +861,7 @@ class TestDataValidation:
861
861
  "period": [0, 1],
862
862
  },
863
863
  )
864
- with pytest.raises(ValueError, match="reserved"):
864
+ with pytest.raises(ValueError, match="collide"):
865
865
  tsam_xarray.aggregate(
866
866
  da,
867
867
  time_dim="time",
@@ -869,6 +869,102 @@ class TestDataValidation:
869
869
  n_clusters=4,
870
870
  )
871
871
 
872
+
873
+ class TestDimNames:
874
+ def _period_da(self) -> xr.DataArray:
875
+ time = pd.date_range("2020-01-01", periods=30 * 24, freq="h")
876
+ return xr.DataArray(
877
+ np.random.default_rng(42).random((len(time), 2, 2)),
878
+ dims=["time", "variable", "period"],
879
+ coords={"time": time, "variable": ["a", "b"], "period": [0, 1]},
880
+ )
881
+
882
+ def test_default_reproduces_current_names(self):
883
+ da = _make_da()
884
+ result = tsam_xarray.aggregate(
885
+ da, n_clusters=4, time_dim="time", cluster_dim="variable"
886
+ )
887
+ assert result.dim_names == tsam_xarray.DimNames()
888
+ assert set(result.cluster_representatives.dims) >= {"cluster", "timestep"}
889
+ assert set(result.cluster_assignments.dims) == {"period", "region"}
890
+
891
+ def test_custom_names_applied(self):
892
+ da = _make_da()
893
+ dn = tsam_xarray.DimNames(cluster="rep", timestep="intra", period="orig_period")
894
+ result = tsam_xarray.aggregate(
895
+ da, n_clusters=4, time_dim="time", cluster_dim="variable", dim_names=dn
896
+ )
897
+ assert set(result.cluster_representatives.dims) >= {"rep", "intra"}
898
+ assert "orig_period" in result.cluster_assignments.dims
899
+ assert set(result.cluster_weights.dims) == {"rep", "region"}
900
+ assert result.n_clusters == 4
901
+ assert result.n_timesteps_per_period == 24
902
+
903
+ def test_renamed_period_frees_input_dim(self):
904
+ da = self._period_da()
905
+ dn = tsam_xarray.DimNames(period="orig_period")
906
+ result = tsam_xarray.aggregate(
907
+ da, n_clusters=4, time_dim="time", cluster_dim="variable", dim_names=dn
908
+ )
909
+ assert "period" in result.cluster_representatives.dims
910
+ assert "orig_period" in result.cluster_assignments.dims
911
+
912
+ def test_collision_still_raises_with_partial_rename(self):
913
+ da = self._period_da()
914
+ dn = tsam_xarray.DimNames(cluster="rep")
915
+ with pytest.raises(ValueError, match="collide"):
916
+ tsam_xarray.aggregate(
917
+ da,
918
+ n_clusters=4,
919
+ time_dim="time",
920
+ cluster_dim="variable",
921
+ dim_names=dn,
922
+ )
923
+
924
+ def test_dim_names_must_be_unique(self):
925
+ with pytest.raises(ValueError, match="unique"):
926
+ tsam_xarray.DimNames(cluster="x", timestep="x")
927
+
928
+ def test_apply_reuses_stored_names(self):
929
+ da = _make_da()
930
+ dn = tsam_xarray.DimNames(cluster="rep", timestep="intra")
931
+ result = tsam_xarray.aggregate(
932
+ da, n_clusters=4, time_dim="time", cluster_dim="variable", dim_names=dn
933
+ )
934
+ applied = result.clustering.apply(da)
935
+ assert set(applied.cluster_representatives.dims) >= {"rep", "intra"}
936
+
937
+ def test_disaggregate_roundtrip_custom_names(self):
938
+ da = _make_da()
939
+ dn = tsam_xarray.DimNames(cluster="rep", timestep="intra")
940
+ result = tsam_xarray.aggregate(
941
+ da, n_clusters=4, time_dim="time", cluster_dim="variable", dim_names=dn
942
+ )
943
+ back = result.disaggregate(result.cluster_representatives)
944
+ assert "time" in back.dims
945
+
946
+ def test_json_roundtrip_preserves_names(self, tmp_path):
947
+ da = _make_da()
948
+ dn = tsam_xarray.DimNames(period="orig_period", cluster="rep")
949
+ result = tsam_xarray.aggregate(
950
+ da, n_clusters=4, time_dim="time", cluster_dim="variable", dim_names=dn
951
+ )
952
+ path = tmp_path / "clustering.json"
953
+ result.clustering.to_json(path)
954
+ loaded = tsam_xarray.ClusteringResult.from_json(path)
955
+ assert loaded.dim_names == dn
956
+ assert "orig_period" in loaded.cluster_assignments.dims
957
+
958
+ def test_from_dict_defaults_names_for_legacy_blob(self):
959
+ da = _make_da()
960
+ result = tsam_xarray.aggregate(
961
+ da, n_clusters=4, time_dim="time", cluster_dim="variable"
962
+ )
963
+ blob = result.clustering.to_dict()
964
+ del blob["dim_names"]
965
+ loaded = tsam_xarray.ClusteringResult.from_dict(blob)
966
+ assert loaded.dim_names == tsam_xarray.DimNames()
967
+
872
968
  def test_dask_array_warns(self):
873
969
  pytest.importorskip("dask")
874
970
  da = _make_da()
@@ -1,3 +0,0 @@
1
- {
2
- ".": "0.6.3"
3
- }
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes
File without changes