openstef-core 4.0.1__tar.gz → 4.1.1.dev0__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 (54) hide show
  1. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/.gitignore +5 -4
  2. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/PKG-INFO +1 -1
  3. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/pyproject.toml +4 -8
  4. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/src/openstef_core/base_model.py +5 -2
  5. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/src/openstef_core/datasets/timeseries_dataset.py +6 -6
  6. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/src/openstef_core/datasets/validated_datasets.py +8 -6
  7. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/src/openstef_core/datasets/validation.py +1 -3
  8. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/src/openstef_core/datasets/versioned_timeseries_dataset.py +3 -6
  9. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/src/openstef_core/exceptions.py +13 -8
  10. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/src/openstef_core/mixins/predictor.py +1 -1
  11. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/src/openstef_core/mixins/stateful.py +4 -3
  12. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/src/openstef_core/mixins/transform.py +1 -0
  13. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/src/openstef_core/testing.py +6 -6
  14. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/src/openstef_core/types.py +15 -13
  15. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/src/openstef_core/utils/__init__.py +6 -0
  16. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/src/openstef_core/utils/multiprocessing.py +2 -4
  17. openstef_core-4.1.1.dev0/src/openstef_core/utils/numpy.py +105 -0
  18. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/src/openstef_core/utils/pandas.py +4 -4
  19. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/tests/unit/datasets/test_mixins.py +14 -10
  20. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/tests/unit/datasets/test_timeseries_dataset.py +45 -35
  21. openstef_core-4.1.1.dev0/tests/unit/datasets/test_validated_datasets.py +382 -0
  22. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/tests/unit/datasets/test_versioned_timeseries_dataset.py +135 -65
  23. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/tests/unit/mixins/test_stateful.py +1 -1
  24. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/tests/unit/mixins/test_transform.py +1 -1
  25. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/tests/unit/test_hyperparams_tuning.py +13 -13
  26. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/tests/unit/test_param_ranges.py +1 -1
  27. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/tests/unit/test_types.py +51 -41
  28. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/tests/unit/utils/test_datetime.py +1 -1
  29. openstef_core-4.1.1.dev0/tests/unit/utils/test_numpy.py +179 -0
  30. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/README.md +0 -0
  31. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/src/openstef_core/__init__.py +0 -0
  32. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/src/openstef_core/constants.py +0 -0
  33. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/src/openstef_core/datasets/__init__.py +0 -0
  34. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/src/openstef_core/datasets/mixins.py +0 -0
  35. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/src/openstef_core/mixins/__init__.py +0 -0
  36. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/src/openstef_core/mixins/param_ranges.py +0 -0
  37. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/src/openstef_core/transforms/__init__.py +0 -0
  38. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/src/openstef_core/transforms/dataset_transforms.py +0 -0
  39. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/src/openstef_core/utils/datetime.py +0 -0
  40. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/src/openstef_core/utils/invariants.py +0 -0
  41. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/src/openstef_core/utils/itertools.py +0 -0
  42. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/src/openstef_core/utils/pydantic.py +0 -0
  43. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/tests/__init__.py +0 -0
  44. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/tests/unit/__init__.py +0 -0
  45. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/tests/unit/datasets/__init__.py +0 -0
  46. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/tests/unit/datasets/test_validation.py +0 -0
  47. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/tests/unit/datasets/utils.py +0 -0
  48. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/tests/unit/mixins/__init__.py +0 -0
  49. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/tests/unit/test_base_model.py +0 -0
  50. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/tests/unit/transforms/__init__.py +0 -0
  51. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/tests/unit/utils/__init__.py +0 -0
  52. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/tests/unit/utils/test_itertools.py +0 -0
  53. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/tests/unit/utils/test_multiprocessing.py +0 -0
  54. {openstef_core-4.0.1 → openstef_core-4.1.1.dev0}/tests/unit/utils/test_pandas.py +0 -0
@@ -35,10 +35,6 @@ MANIFEST
35
35
  # Ruff
36
36
  .ruff_cache/
37
37
 
38
- # Pyright
39
- .pyright/
40
- # pyright-report/
41
-
42
38
  # Test, coverage, tox
43
39
  .pytest_cache/
44
40
  .coverage
@@ -69,6 +65,8 @@ docs/_build/
69
65
  docs/source/api/generated/
70
66
  docs/source/tutorials/
71
67
  docs/source/benchmarks/
68
+ # Community health files materialized from OpenSTEF/.github at build time
69
+ docs/source/contribute/_community/
72
70
  docs/source/user_guide/**/quick_start_tutorial.py
73
71
  docs/source/user_guide/**/feature_engineering_tutorial.py
74
72
  docs/source/user_guide/**/datasets_tutorial.py
@@ -135,6 +133,9 @@ benchmark_results*/
135
133
  # Local dataset files
136
134
  liander_dataset/
137
135
 
136
+ # Deployment example run artifacts (MLflow store, forecasts, dataset, Celery/Airflow state)
137
+ openstef_deployment_runs/
138
+
138
139
  # Mlflow
139
140
  /mlflow
140
141
  /mlflow_artifacts_local
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: openstef-core
3
- Version: 4.0.1
3
+ Version: 4.1.1.dev0
4
4
  Summary: Core functionality for OpenSTEF, a framework for short-term energy forecasting.
5
5
  Project-URL: Documentation, https://openstef.github.io/openstef/index.html
6
6
  Project-URL: Homepage, https://lfenergy.org/projects/openstef/
@@ -1,15 +1,13 @@
1
1
  # SPDX-FileCopyrightText: 2025 Contributors to the OpenSTEF project <openstef@lfenergy.org>
2
2
  #
3
3
  # SPDX-License-Identifier: MPL-2.0
4
-
5
4
  [build-system]
6
5
  build-backend = "hatchling.build"
7
-
8
6
  requires = [ "hatchling" ]
9
7
 
10
8
  [project]
11
9
  name = "openstef-core"
12
- version = "4.0.1"
10
+ version = "4.1.1.dev0"
13
11
  description = "Core functionality for OpenSTEF, a framework for short-term energy forecasting."
14
12
  readme = "README.md"
15
13
  keywords = [ "energy", "forecasting", "machinelearning" ]
@@ -26,24 +24,22 @@ classifiers = [
26
24
  "Programming Language :: Python :: 3.13",
27
25
  "Programming Language :: Python :: 3.14",
28
26
  ]
29
-
30
27
  dependencies = [
31
28
  "joblib>=1,<2",
32
29
  "numpy>=2.3.2,<3",
30
+ # Held at <3 pending the pandas 3.0 Copy-on-Write migration (see migration issue).
33
31
  "pandas>=2.3.1,<3",
34
32
  "pyarrow>=21",
35
33
  "pydantic>=2.12.4,<3",
36
34
  "pydantic-extra-types>=2.10.5,<3",
37
35
  ]
38
-
39
36
  optional-dependencies.benchmark = [
40
37
  "huggingface-hub>=1.2.2",
41
38
  ]
42
-
43
39
  urls.Documentation = "https://openstef.github.io/openstef/index.html"
44
40
  urls.Homepage = "https://lfenergy.org/projects/openstef/"
45
41
  urls.Issues = "https://github.com/OpenSTEF/openstef/issues"
46
42
  urls.Repository = "https://github.com/OpenSTEF/openstef"
47
43
 
48
- [tool.hatch.build.targets.wheel]
49
- packages = [ "src/openstef_core" ]
44
+ [tool.hatch]
45
+ build.targets.wheel.packages = [ "src/openstef_core" ]
@@ -11,7 +11,7 @@ operate on arbitrary config instances or Pydantic models / adapters.
11
11
  """
12
12
 
13
13
  from pathlib import Path
14
- from typing import Annotated, Any, Self
14
+ from typing import Annotated, Any, Self, override
15
15
 
16
16
  import yaml
17
17
  from pydantic import BaseModel as PydanticBaseModel
@@ -112,6 +112,7 @@ def read_yaml_config[T: BaseConfig, U](path: Path, class_type: type[T] | TypeAda
112
112
  class PydanticStringPrimitive:
113
113
  """Base class for Pydantic-compatible types with string serialization."""
114
114
 
115
+ @override
115
116
  def __str__(self) -> str:
116
117
  """Convert to string representation."""
117
118
  raise NotImplementedError("Subclasses must implement __str__")
@@ -158,7 +159,7 @@ class PydanticStringPrimitive:
158
159
  )
159
160
 
160
161
  @classmethod
161
- def __get_pydantic_json_schema__( # noqa: PLW3201
162
+ def __get_pydantic_json_schema__(
162
163
  cls,
163
164
  _schema: core_schema.CoreSchema,
164
165
  handler: GetJsonSchemaHandler,
@@ -172,6 +173,7 @@ class PydanticStringPrimitive:
172
173
  """
173
174
  return {"type": "string"}
174
175
 
176
+ @override
175
177
  def __eq__(self, other: object) -> bool:
176
178
  """Check equality based on string representation.
177
179
 
@@ -182,6 +184,7 @@ class PydanticStringPrimitive:
182
184
  return NotImplemented
183
185
  return str(self) == str(other)
184
186
 
187
+ @override
185
188
  def __hash__(self) -> int:
186
189
  """Return hash based on string representation."""
187
190
  return hash(str(self))
@@ -33,7 +33,7 @@ from openstef_core.utils.pandas import unsafe_sorted_range_slice_idxs
33
33
  _logger = logging.getLogger(__name__)
34
34
 
35
35
 
36
- class TimeSeriesDataset(TimeSeriesMixin, DatasetMixin): # noqa: PLR0904 - important utility class, allow too many public methods
36
+ class TimeSeriesDataset(TimeSeriesMixin, DatasetMixin):
37
37
  """A time series dataset with regular sampling intervals and optional versioning.
38
38
 
39
39
  This class represents time series data with a consistent sampling interval
@@ -304,7 +304,7 @@ class TimeSeriesDataset(TimeSeriesMixin, DatasetMixin): # noqa: PLR0904 - impor
304
304
  Returns:
305
305
  New dataset containing only rows with timestamps in the mask.
306
306
  """
307
- data_filtered = self.data.loc[self.index.isin(mask)] # pyright: ignore[reportUnknownMemberType]
307
+ data_filtered = self.data.loc[self.index.isin(mask)]
308
308
 
309
309
  return self._copy_with_data(data=data_filtered)
310
310
 
@@ -320,7 +320,7 @@ class TimeSeriesDataset(TimeSeriesMixin, DatasetMixin): # noqa: PLR0904 - impor
320
320
  if self.horizons is None:
321
321
  return self
322
322
 
323
- data_selected = self.data[self.lead_time_series == horizon.value]
323
+ data_selected = cast(pd.DataFrame, self.data[self.lead_time_series == horizon.value])
324
324
  return self._copy_with_data(data=data_selected)
325
325
 
326
326
  def to_pandas(self) -> pd.DataFrame:
@@ -395,7 +395,7 @@ class TimeSeriesDataset(TimeSeriesMixin, DatasetMixin): # noqa: PLR0904 - impor
395
395
  available_at_column: str = "available_at",
396
396
  horizon_column: str = "horizon",
397
397
  ) -> Self:
398
- df = pd.read_parquet(path=path) # pyright: ignore[reportUnknownMemberType]
398
+ df = pd.read_parquet(path=path)
399
399
  if not isinstance(df.index, pd.DatetimeIndex):
400
400
  if timestamp_column not in df.columns:
401
401
  raise TimeSeriesValidationError(
@@ -512,8 +512,8 @@ def validate_horizons_present(dataset: TimeSeriesDataset, horizons: list[LeadTim
512
512
  if dataset.horizons is None and len(horizons) == 1:
513
513
  return # Non-versioned dataset can satisfy single-horizon requests
514
514
 
515
- required_horizons = set(horizons or [])
516
- missing_horizons = [h for h in horizons if h not in required_horizons]
515
+ available_horizons = set(dataset.horizons or [])
516
+ missing_horizons = [h for h in horizons if h not in available_horizons]
517
517
  if missing_horizons:
518
518
  raise TimeSeriesValidationError("Missing forecast horizons: " + ", ".join(map(str, missing_horizons)))
519
519
 
@@ -345,7 +345,7 @@ class ForecastDataset(TimeSeriesDataset):
345
345
  """
346
346
  if self.standard_deviation_column not in self.data.columns:
347
347
  raise MissingColumnsError(missing_columns=[self.standard_deviation_column])
348
- return self.data[self.standard_deviation_column] # pyright: ignore[reportUnknownVariableType]
348
+ return self.data[self.standard_deviation_column]
349
349
 
350
350
  @property
351
351
  def quantiles_data(self) -> pd.DataFrame:
@@ -572,11 +572,13 @@ class EnsembleForecastDataset(TimeSeriesDataset):
572
572
  if sample_weights is not None:
573
573
  additional_columns[sample_weight_column] = sample_weights
574
574
 
575
- combined_data = pd.DataFrame({
576
- f"{learner}{ENSEMBLE_COLUMN_SEP}{q.format()}": ds.data[q.format()]
577
- for learner, ds in datasets.items()
578
- for q in ds.quantiles
579
- }).assign(**additional_columns)
575
+ combined_data = pd.DataFrame(
576
+ {
577
+ f"{learner}{ENSEMBLE_COLUMN_SEP}{q.format()}": ds.data[q.format()]
578
+ for learner, ds in datasets.items()
579
+ for q in ds.quantiles
580
+ }
581
+ ).assign(**additional_columns)
580
582
 
581
583
  return cls(
582
584
  data=combined_data,
@@ -8,8 +8,6 @@ This module provides functions to validate dataset compatibility and integrity,
8
8
  particularly for operations that combine multiple datasets.
9
9
  """
10
10
 
11
- import functools
12
- import operator
13
11
  from collections import Counter
14
12
  from collections.abc import Iterable, Sequence
15
13
  from datetime import timedelta
@@ -53,7 +51,7 @@ def validate_disjoint_columns(datasets: Iterable[TimeSeriesMixin]) -> list[str]:
53
51
  Raises:
54
52
  TimeSeriesValidationError: If any feature name appears in multiple datasets.
55
53
  """
56
- all_features: list[str] = functools.reduce(operator.iadd, [d.feature_names for d in datasets], [])
54
+ all_features: list[str] = [feature for d in datasets for feature in d.feature_names]
57
55
  if len(all_features) != len(set(all_features)):
58
56
  duplicate_features = [item for item, count in Counter(all_features).items() if count > 1]
59
57
  raise TimeSeriesValidationError("Datasets have overlapping feature names: " + ", ".join(duplicate_features))
@@ -182,7 +182,7 @@ class VersionedTimeSeriesDataset(TimeSeriesMixin, DatasetMixin):
182
182
  available_at_column: str = "available_at",
183
183
  horizon_column: str = "horizon",
184
184
  ) -> Self:
185
- df = pd.read_parquet(path=path) # type: ignore
185
+ df = pd.read_parquet(path=path)
186
186
  if "parts" in df.attrs:
187
187
  parts_metadata = json.loads(df.attrs.get("parts", "{}")).get("parts", [])
188
188
  if len(parts_metadata) == 0:
@@ -190,7 +190,7 @@ class VersionedTimeSeriesDataset(TimeSeriesMixin, DatasetMixin):
190
190
 
191
191
  parts: list[TimeSeriesDataset] = [
192
192
  TimeSeriesDataset(
193
- data=df.loc[df.part_id == i, part_info["columns"]],
193
+ data=cast("pd.DataFrame", df.loc[df.part_id == i, part_info["columns"]]),
194
194
  sample_interval=timedelta_from_isoformat(part_info.get("sample_interval", "PT1H")),
195
195
  )
196
196
  for i, part_info in enumerate(parts_metadata)
@@ -251,10 +251,7 @@ class VersionedTimeSeriesDataset(TimeSeriesMixin, DatasetMixin):
251
251
  index = functools.reduce(lambda x, y: x.intersection(y), [part.index.unique() for part in data_parts])
252
252
 
253
253
  return cls(
254
- data_parts=[
255
- TimeSeriesDataset(data=part.data.loc[part.index.isin(index)]) # pyright: ignore[reportUnknownMemberType]
256
- for part in data_parts
257
- ],
254
+ data_parts=[TimeSeriesDataset(data=part.data.loc[part.index.isin(index)]) for part in data_parts],
258
255
  index=index,
259
256
  )
260
257
 
@@ -12,19 +12,24 @@ from collections.abc import Sequence
12
12
 
13
13
 
14
14
  class MissingExtraError(Exception):
15
- """Exception raised when an extra is missing in the extras list."""
15
+ """Exception raised when an optional dependency for a feature is not installed."""
16
16
 
17
- def __init__(self, extra: str, package: str = "openstef-beam"):
18
- """Initialize the exception with the name of the missing extra.
17
+ def __init__(self, extra: str, package: str = "openstef-beam", *, install_extra: str | None = None):
18
+ """Initialize the exception with the missing dependency and how to install it.
19
19
 
20
20
  Args:
21
- extra: Name of the missing extra package.
22
- package: Name of the package requiring the extra.
21
+ extra: Import/distribution name that failed to import (e.g. ``"onnxruntime"``).
22
+ package: OpenSTEF package that owns the optional feature.
23
+ install_extra: The extra group to install for this feature (e.g. ``"cpu"``),
24
+ yielding ``pip install {package}[{install_extra}]``. When ``None`` the hint
25
+ installs ``extra`` directly. Prefer naming the extra group: not every package
26
+ defines an ``[all]`` extra, and some groups are mutually exclusive.
23
27
  """
24
28
  self.extra = extra
29
+ target = f"{package}[{install_extra}]" if install_extra else extra
25
30
  super().__init__(
26
- f"Optional package {extra} is missing. Please install it to use this module using `pip install {extra}` "
27
- f"or install all optional features using `pip install {package}[all]`."
31
+ f"Optional dependency '{extra}' is required for this feature but is not installed. "
32
+ f"Install it with `pip install {target}`."
28
33
  )
29
34
 
30
35
 
@@ -100,7 +105,7 @@ class ModelUnderperformingError(Exception):
100
105
  threshold: The defined threshold for acceptable performance.
101
106
  """
102
107
  message = (
103
- f"Model is underperforming: {metric_name} = {metric_value:.4f}"
108
+ f"Model is underperforming: {metric_name} = {metric_value:.4f} "
104
109
  f"does not meet the threshold of {threshold:.4f}."
105
110
  )
106
111
  super().__init__(message)
@@ -234,7 +234,7 @@ class HyperParams(BaseConfig):
234
234
  result: HyperParams = handler(data)
235
235
  if instance_ranges and result.__pydantic_private__ is not None:
236
236
  result._instance_ranges = instance_ranges
237
- return result # type: ignore[return-value]
237
+ return result # ty: ignore[invalid-return-type]
238
238
 
239
239
  def get_search_space(self, include: set[str] | None = None) -> dict[str, TuningRange]:
240
240
  """Merge instance and class-level ranges, returning only ``tune=True`` fields.
@@ -10,7 +10,7 @@ serialization with automatic state migration.
10
10
  """
11
11
 
12
12
  import warnings
13
- from typing import ClassVar, TypedDict, cast
13
+ from typing import ClassVar, TypedDict, cast, override
14
14
 
15
15
  from openstef_core.types import Any
16
16
 
@@ -41,6 +41,7 @@ class Stateful:
41
41
 
42
42
  _VERSION: ClassVar[int] = 1
43
43
 
44
+ @override
44
45
  def __getstate__(self) -> VersionedState:
45
46
  """Serialize object state with version metadata.
46
47
 
@@ -74,7 +75,7 @@ class Stateful:
74
75
  state: Serialized state, either VersionedState dict or legacy format.
75
76
  """
76
77
  # Handle legacy objects without versioning
77
- if not isinstance(state, dict) or "__version__" not in state: # pyright: ignore[reportUnnecessaryIsInstance]
78
+ if not isinstance(state, dict) or "__version__" not in state:
78
79
  warnings.warn(
79
80
  f"Loading legacy {self.__class__.__name__} without version metadata.", UserWarning, stacklevel=2
80
81
  )
@@ -108,7 +109,7 @@ class Stateful:
108
109
  """
109
110
  # Check if any parent class has __setstate__
110
111
  if hasattr(super(), "__setstate__"):
111
- super().__setstate__(state) # type: ignore[misc]
112
+ super().__setstate__(state) # ty: ignore[unresolved-attribute]
112
113
  elif state: # Only update if state is not empty
113
114
  self.__dict__.update(state)
114
115
 
@@ -130,6 +130,7 @@ class TransformPipeline[T](BaseModel, Transform[T, T]):
130
130
  description="Sequence of transforms to apply in sequence. If empty, the pipeline is a nop.",
131
131
  )
132
132
 
133
+ @override
133
134
  def __reduce__(self) -> tuple[Callable[[], "TransformPipeline[Any]"], tuple[()], Any]:
134
135
  """Support pickling of generic TransformPipeline instances.
135
136
 
@@ -78,7 +78,7 @@ def create_timeseries_dataset(
78
78
  )
79
79
 
80
80
 
81
- def create_synthetic_forecasting_dataset( # noqa: PLR0913, PLR0917 - complex function - testing utility
81
+ def create_synthetic_forecasting_dataset( # noqa: PLR0913 - complex function - testing utility
82
82
  start: datetime = datetime.fromisoformat("2025-01-01T00:00:00+00:00"), # noqa: B008
83
83
  length: timedelta = timedelta(days=30 * 9),
84
84
  sample_interval: timedelta = timedelta(hours=1),
@@ -169,9 +169,6 @@ def load_liander_dataset(
169
169
  Downloads load measurements, weather forecasts, electricity prices, and standard load
170
170
  profiles from HuggingFace Hub, then combines them via left join.
171
171
 
172
- Raises:
173
- ImportError: When ``huggingface-hub`` is not installed.
174
-
175
172
  Args:
176
173
  target: Sub-path within the repo identifying the installation (e.g. ``"mv_feeder/OS Gorredijk"``).
177
174
  repo_id: HuggingFace dataset repository ID.
@@ -180,9 +177,12 @@ def load_liander_dataset(
180
177
 
181
178
  Returns:
182
179
  Combined dataset with all features aligned by timestamp.
180
+
181
+ Raises:
182
+ ImportError: When ``huggingface-hub`` is not installed.
183
183
  """
184
184
  try:
185
- from huggingface_hub import hf_hub_download # pyright: ignore[reportUnknownVariableType] # noqa: PLC0415
185
+ from huggingface_hub import hf_hub_download # noqa: PLC0415
186
186
  from huggingface_hub.utils import logging as hf_logging # noqa: PLC0415
187
187
  except ImportError:
188
188
  msg = "huggingface-hub is required for benchmark datasets: pip install openstef-core[benchmark]"
@@ -199,7 +199,7 @@ def load_liander_dataset(
199
199
  # Suppress HF Hub noise (unauthenticated requests warning, progress bars)
200
200
  hf_logging.set_verbosity_error()
201
201
  for filename in files_to_download:
202
- hf_hub_download( # pyright: ignore[reportCallIssue]
202
+ hf_hub_download(
203
203
  repo_id=repo_id,
204
204
  filename=filename,
205
205
  repo_type="dataset",
@@ -18,9 +18,9 @@ from decimal import Decimal
18
18
  from enum import StrEnum
19
19
  from functools import total_ordering
20
20
  from typing import Any, Literal, Self, override
21
+ from zoneinfo import ZoneInfo
21
22
 
22
23
  import pandas as pd
23
- import pytz
24
24
  from pydantic import GetCoreSchemaHandler, TypeAdapter, ValidationInfo
25
25
  from pydantic_core import CoreSchema, core_schema
26
26
 
@@ -55,6 +55,7 @@ class LeadTime(PydanticStringPrimitive):
55
55
  """
56
56
  self.value = value
57
57
 
58
+ @override
58
59
  def __str__(self) -> str:
59
60
  """Converts to ISO 8601 duration string.
60
61
 
@@ -63,6 +64,7 @@ class LeadTime(PydanticStringPrimitive):
63
64
  """
64
65
  return TypeAdapter(timedelta).dump_python(self.value, mode="json")
65
66
 
67
+ @override
66
68
  def __repr__(self) -> str:
67
69
  """Returns a detailed string representation for debugging.
68
70
 
@@ -72,6 +74,7 @@ class LeadTime(PydanticStringPrimitive):
72
74
  return f"LeadTime('{self}')"
73
75
 
74
76
  @classmethod
77
+ @override
75
78
  def from_string(cls, s: str) -> Self:
76
79
  """Creates an instance from an ISO 8601 duration string.
77
80
 
@@ -131,10 +134,10 @@ class AvailableAt(PydanticStringPrimitive):
131
134
  - *HHMM* is the time of day
132
135
 
133
136
  An optional timezone suffix ``[Region/City]`` (RFC 9557 bracket
134
- notation) makes the availability time timezone-aware. Both pytz
137
+ notation) makes the availability time timezone-aware. ``zoneinfo``
135
138
  and stdlib ``datetime.timezone`` objects are accepted; they
136
139
  round-trip through the IANA name via ``str(tz)`` /
137
- ``pytz.timezone(name)``.
140
+ ``zoneinfo.ZoneInfo(name)``.
138
141
 
139
142
  For example, ``D-1T0600[Europe/Amsterdam]`` means "6:00
140
143
  Europe/Amsterdam on the previous day".
@@ -143,8 +146,8 @@ class AvailableAt(PydanticStringPrimitive):
143
146
 
144
147
  Example:
145
148
  >>> from datetime import time
146
- >>> import pytz
147
- >>> tz_at = AvailableAt(day_offset=-1, time_of_day=time(6, 0), tzinfo=pytz.timezone('Europe/Amsterdam'))
149
+ >>> from zoneinfo import ZoneInfo
150
+ >>> tz_at = AvailableAt(day_offset=-1, time_of_day=time(6, 0), tzinfo=ZoneInfo('Europe/Amsterdam'))
148
151
  >>> str(tz_at)
149
152
  'D-1T0600[Europe/Amsterdam]'
150
153
  >>> at = AvailableAt.from_string("D-1T0600")
@@ -152,7 +155,7 @@ class AvailableAt(PydanticStringPrimitive):
152
155
  (-1, datetime.time(6, 0))
153
156
  """
154
157
 
155
- def __init__(self, day_offset: int, time_of_day: time, *, tzinfo: pytz.BaseTzInfo | dt_timezone | None = None):
158
+ def __init__(self, day_offset: int, time_of_day: time, *, tzinfo: ZoneInfo | dt_timezone | None = None):
156
159
  """Initialise with a day offset, time of day, and optional timezone.
157
160
 
158
161
  Args:
@@ -160,7 +163,7 @@ class AvailableAt(PydanticStringPrimitive):
160
163
  ``-1`` means "the previous day", ``0`` means "the same day".
161
164
  time_of_day: Clock time when data becomes available.
162
165
  tzinfo: Optional timezone for the availability time
163
- (e.g. ``pytz.timezone("Europe/Amsterdam")``, ``pytz.UTC``,
166
+ (e.g. ``zoneinfo.ZoneInfo("Europe/Amsterdam")``, ``zoneinfo.ZoneInfo("UTC")``,
164
167
  or ``datetime.timezone.utc``).
165
168
 
166
169
  Raises:
@@ -173,6 +176,7 @@ class AvailableAt(PydanticStringPrimitive):
173
176
  self.time_of_day = time_of_day
174
177
  self.tzinfo = tzinfo
175
178
 
179
+ @override
176
180
  def __str__(self) -> str:
177
181
  """Converts to string in ``DnTHHMM`` or ``DnTHHMM[tz]`` format.
178
182
 
@@ -185,6 +189,7 @@ class AvailableAt(PydanticStringPrimitive):
185
189
  return base
186
190
 
187
191
  @classmethod
192
+ @override
188
193
  def from_string(cls, s: str) -> Self:
189
194
  """Creates an instance from a string in ``DnTHHMM[tz]`` format.
190
195
 
@@ -212,9 +217,9 @@ class AvailableAt(PydanticStringPrimitive):
212
217
  raise ValueError(msg)
213
218
 
214
219
  if z_part:
215
- resolved_tz = pytz.UTC
220
+ resolved_tz = ZoneInfo("UTC")
216
221
  elif tz_part:
217
- resolved_tz = pytz.timezone(tz_part)
222
+ resolved_tz = ZoneInfo(tz_part)
218
223
  else:
219
224
  resolved_tz = None
220
225
 
@@ -256,10 +261,7 @@ class AvailableAt(PydanticStringPrimitive):
256
261
  if source_tz is None:
257
262
  return naive_result
258
263
 
259
- if isinstance(source_tz, pytz.BaseTzInfo):
260
- aware = source_tz.localize(naive_result)
261
- else:
262
- aware = naive_result.replace(tzinfo=source_tz)
264
+ aware = naive_result.replace(tzinfo=source_tz)
263
265
 
264
266
  if date.tzinfo is not None:
265
267
  return aware.astimezone(date.tzinfo)
@@ -19,6 +19,10 @@ from openstef_core.utils.invariants import (
19
19
  from openstef_core.utils.multiprocessing import (
20
20
  run_parallel,
21
21
  )
22
+ from openstef_core.utils.numpy import (
23
+ interpolate_quantiles,
24
+ zero_fill_with_mask,
25
+ )
22
26
  from openstef_core.utils.pydantic import (
23
27
  timedelta_from_isoformat,
24
28
  timedelta_to_isoformat,
@@ -27,8 +31,10 @@ from openstef_core.utils.pydantic import (
27
31
  __all__ = [
28
32
  "align_datetime",
29
33
  "align_datetime_to_time",
34
+ "interpolate_quantiles",
30
35
  "not_none",
31
36
  "run_parallel",
32
37
  "timedelta_from_isoformat",
33
38
  "timedelta_to_isoformat",
39
+ "zero_fill_with_mask",
34
40
  ]
@@ -63,12 +63,10 @@ def run_parallel[T, R](
63
63
  return [process_fn(item) for item in items]
64
64
 
65
65
  if mode == "loky":
66
- from joblib import Parallel, delayed # pyright: ignore[reportUnknownVariableType] # noqa: PLC0415
66
+ from joblib import Parallel, delayed # noqa: PLC0415
67
67
 
68
68
  # Use joblib with loky backend for robust process management
69
- return Parallel(n_jobs=n_processes, backend="loky")( # pyright: ignore[reportUnknownVariableType]
70
- delayed(process_fn)(item) for item in items
71
- ) # type: ignore
69
+ return Parallel(n_jobs=n_processes, backend="loky")(delayed(process_fn)(item) for item in items)
72
70
 
73
71
  # Auto-configure for macOS
74
72
  context = multiprocessing.get_context(method=mode)
@@ -0,0 +1,105 @@
1
+ # SPDX-FileCopyrightText: 2025 Contributors to the OpenSTEF project <openstef@lfenergy.org>
2
+ #
3
+ # SPDX-License-Identifier: MPL-2.0
4
+
5
+ """Pure NumPy helpers for probabilistic forecasting."""
6
+
7
+ import logging
8
+ from collections.abc import Sequence
9
+
10
+ import numpy as np
11
+
12
+ logger = logging.getLogger(__name__)
13
+
14
+
15
+ def zero_fill_with_mask(values: np.ndarray) -> tuple[np.ndarray, np.ndarray]:
16
+ """Split an array into a zero-filled copy and a finiteness mask.
17
+
18
+ Non-finite entries (``NaN`` or infinities) are replaced with ``0.0`` in the
19
+ returned values; the mask is ``1.0`` exactly where the original entry was
20
+ finite. This is the common "feed raw values, tell the model which are real"
21
+ step for masked model inputs.
22
+
23
+ Args:
24
+ values: Array of any shape.
25
+
26
+ Returns:
27
+ Tuple ``(filled, mask)``, both ``float32`` and the same shape as
28
+ *values*. ``filled`` has non-finite entries zeroed; ``mask`` is ``1.0``
29
+ where *values* was finite, ``0.0`` otherwise.
30
+ """
31
+ finite = np.isfinite(values)
32
+ filled = np.where(finite, values, np.float32(0.0)).astype(np.float32)
33
+ return filled, finite.astype(np.float32)
34
+
35
+
36
+ def interpolate_quantiles(
37
+ predictions: np.ndarray,
38
+ source_quantiles: Sequence[float],
39
+ target_quantiles: Sequence[float],
40
+ ) -> np.ndarray:
41
+ """Resample quantile predictions onto a new quantile grid.
42
+
43
+ Performs piecewise-linear interpolation across the quantile dimension (the
44
+ last axis of *predictions*). Target levels outside the source range are
45
+ clamped to the nearest source prediction (constant extrapolation), which
46
+ keeps the resampled values within the predicted envelope. Because the model
47
+ cannot say anything beyond its most extreme level, a request for, say, q0.999
48
+ against a model whose highest level is q0.99 returns the q0.99 prediction; a
49
+ warning is logged whenever this clamping happens.
50
+
51
+ Args:
52
+ predictions: Array of shape ``(..., n_source)`` whose last axis holds
53
+ predictions for each level in *source_quantiles*, in the same order.
54
+ source_quantiles: Strictly ascending quantile levels the model emits.
55
+ target_quantiles: Quantile levels to resample onto. Any order.
56
+
57
+ Returns:
58
+ Array of shape ``(..., n_target)`` with predictions for each level in
59
+ *target_quantiles*, in the same order.
60
+
61
+ Raises:
62
+ ValueError: If *source_quantiles* is not strictly ascending, or its
63
+ length does not match the last axis of *predictions*.
64
+ """
65
+ src = np.asarray(source_quantiles, dtype=np.float64)
66
+ tgt = np.asarray(target_quantiles, dtype=np.float64)
67
+
68
+ min_levels = 2 # need at least two source levels to interpolate between
69
+ if src.ndim != 1 or src.shape[0] < min_levels:
70
+ msg = "source_quantiles must be a 1-D sequence with at least two levels."
71
+ raise ValueError(msg)
72
+ if predictions.shape[-1] != src.shape[0]:
73
+ msg = (
74
+ f"predictions last axis ({predictions.shape[-1]}) must match the number "
75
+ f"of source quantiles ({src.shape[0]})."
76
+ )
77
+ raise ValueError(msg)
78
+ if np.any(np.diff(src) <= 0):
79
+ msg = "source_quantiles must be strictly ascending."
80
+ raise ValueError(msg)
81
+
82
+ out_of_range = tgt[(tgt < src[0]) | (tgt > src[-1])]
83
+ if out_of_range.size:
84
+ logger.warning(
85
+ "Target quantile level(s) %s lie outside the source range [%s, %s]; "
86
+ "clamping to the nearest source quantile (constant extrapolation).",
87
+ np.unique(out_of_range).tolist(),
88
+ src[0],
89
+ src[-1],
90
+ )
91
+
92
+ # Bracket each target level by the adjacent source levels, clamping the
93
+ # endpoints so out-of-range targets extrapolate as constants.
94
+ upper = np.clip(np.searchsorted(src, tgt, side="left"), 1, src.shape[0] - 1)
95
+ lower = upper - 1
96
+
97
+ weight = (tgt - src[lower]) / (src[upper] - src[lower])
98
+ weight = np.clip(weight, 0.0, 1.0)
99
+
100
+ low_values = predictions[..., lower]
101
+ high_values = predictions[..., upper]
102
+ return low_values * (1.0 - weight) + high_values * weight
103
+
104
+
105
+ __all__ = ["interpolate_quantiles", "zero_fill_with_mask"]