diff --git a/packages/openstef-models/src/openstef_models/transforms/general/__init__.py b/packages/openstef-models/src/openstef_models/transforms/general/__init__.py index c7f91b228..509a64026 100644 --- a/packages/openstef-models/src/openstef_models/transforms/general/__init__.py +++ b/packages/openstef-models/src/openstef_models/transforms/general/__init__.py @@ -20,6 +20,7 @@ from openstef_models.transforms.general.scaler import Scaler from openstef_models.transforms.general.selector import Selector from openstef_models.transforms.general.shifter import Shifter +from openstef_models.transforms.general.sklearn_adapter import SklearnTransformAdapter __all__ = [ "OUTLIER_NAN_MASK_PREFIX", @@ -34,4 +35,5 @@ "Scaler", "Selector", "Shifter", + "SklearnTransformAdapter", ] diff --git a/packages/openstef-models/src/openstef_models/transforms/general/sklearn_adapter.py b/packages/openstef-models/src/openstef_models/transforms/general/sklearn_adapter.py new file mode 100644 index 000000000..8234568c9 --- /dev/null +++ b/packages/openstef-models/src/openstef_models/transforms/general/sklearn_adapter.py @@ -0,0 +1,131 @@ +# SPDX-FileCopyrightText: 2025 Contributors to the OpenSTEF project +# +# SPDX-License-Identifier: MPL-2.0 + +"""Adapter that exposes any scikit-learn transformer as a TimeSeriesTransform.""" + +import importlib +from typing import Any, override + +import pandas as pd +from pydantic import Field, PrivateAttr + +from openstef_core.base_model import BaseConfig +from openstef_core.datasets import TimeSeriesDataset +from openstef_core.exceptions import MissingExtraError, NotFittedError +from openstef_core.transforms import TimeSeriesTransform +from openstef_models.utils.feature_selection import FeatureSelection + + +class SklearnTransformAdapter(BaseConfig, TimeSeriesTransform): + """Adapt any scikit-learn transformer to the OpenSTEF ``TimeSeriesTransform`` interface. + + The transformer is specified by its import path and constructor parameters rather than + as an object, so the configuration stays serializable (save/load-able). It is fitted on + the selected feature columns; its output then replaces those columns while the remaining + columns pass through unchanged. Output column names come from the transformer's + ``get_feature_names_out()``, so shape-changing transforms (e.g. PCA, one-hot encoders) + are handled the same way as shape-preserving ones (e.g. scalers). + + ``features_added()`` is populated after ``fit()`` (the output names of some transformers + are only known once fitted). + + Example: + >>> import pandas as pd + >>> from datetime import timedelta + >>> from openstef_core.datasets import TimeSeriesDataset + >>> from openstef_models.transforms.general import SklearnTransformAdapter + >>> + >>> data = pd.DataFrame( + ... {"load": [100.0, 200.0, 300.0]}, + ... index=pd.date_range("2025-01-01", periods=3, freq="h"), + ... ) + >>> dataset = TimeSeriesDataset(data, timedelta(hours=1)) + >>> adapter = SklearnTransformAdapter(transformer_class="sklearn.preprocessing.StandardScaler") + >>> adapter.fit(dataset) + >>> transformed = adapter.transform(dataset) + >>> abs(float(transformed.data["load"].mean().round(6))) + 0.0 + >>> adapter.features_added() + [] + """ + + transformer_class: str = Field( + description="Import path of the scikit-learn transformer, e.g. 'sklearn.preprocessing.StandardScaler'.", + ) + transformer_params: dict[str, Any] = Field( + default_factory=dict, + description="Keyword arguments passed to the transformer's constructor.", + ) + selection: FeatureSelection = Field( + default=FeatureSelection.ALL, + description="Features the transformer is applied to.", + ) + + _transformer: Any = PrivateAttr() + _is_fitted: bool = PrivateAttr(default=False) + _added_features: list[str] = PrivateAttr(default_factory=list) + + @property + @override + def is_fitted(self) -> bool: + return self._is_fitted + + @override + def model_post_init(self, context: Any) -> None: + # Restrict to scikit-learn to keep the dynamic import from loading arbitrary modules. + if not self.transformer_class.startswith("sklearn."): + msg = f"transformer_class must be a scikit-learn transformer (sklearn.*), got {self.transformer_class!r}." + raise ValueError(msg) + module_path, _, class_name = self.transformer_class.rpartition(".") + try: + module = importlib.import_module(module_path) + except ImportError as e: + raise MissingExtraError("sklearn", package="openstef-models") from e + try: + transformer_cls = getattr(module, class_name) + except AttributeError as e: + msg = f"{class_name!r} was not found in {module_path!r}." + raise ValueError(msg) from e + self._transformer = transformer_cls(**self.transformer_params) + + def _output_names(self, features: list[str]) -> list[str]: + # Prefer the transformer's own output names (handles shape-changing transforms); + # fall back to the input names for shape-preserving ones. + if hasattr(self._transformer, "get_feature_names_out"): + return list(self._transformer.get_feature_names_out(features)) + return list(features) + + @override + def fit(self, data: TimeSeriesDataset) -> None: + features = self.selection.resolve(data.feature_names) + self._transformer.fit(data.data[features]) + output_names = self._output_names(features) + self._added_features = [name for name in output_names if name not in data.feature_names] + self._is_fitted = True + + @override + def transform(self, data: TimeSeriesDataset) -> TimeSeriesDataset: + if not self._is_fitted: + raise NotFittedError(self.__class__.__name__) + + features = self.selection.resolve(data.feature_names) + output_names = self._output_names(features) + transformed = pd.DataFrame( + self._transformer.transform(data.data[features]), + index=data.data.index, + columns=output_names, + ) + + # Replace the transformed inputs with the transformer's output, keep the rest. + passthrough = [column for column in data.data.columns if column not in features] + result = pd.concat([data.data[passthrough], transformed], axis=1) + + return TimeSeriesDataset(data=result, sample_interval=data.sample_interval) + + @override + def features_added(self) -> list[str]: + return self._added_features + + +__all__ = ["SklearnTransformAdapter"] diff --git a/packages/openstef-models/tests/unit/transforms/general/test_sklearn_adapter.py b/packages/openstef-models/tests/unit/transforms/general/test_sklearn_adapter.py new file mode 100644 index 000000000..2dcc2349b --- /dev/null +++ b/packages/openstef-models/tests/unit/transforms/general/test_sklearn_adapter.py @@ -0,0 +1,103 @@ +# SPDX-FileCopyrightText: 2025 Contributors to the OpenSTEF project +# +# SPDX-License-Identifier: MPL-2.0 + +from datetime import timedelta + +import pandas as pd +import pytest + +from openstef_core.datasets import TimeSeriesDataset +from openstef_core.exceptions import NotFittedError +from openstef_models.transforms.general import SklearnTransformAdapter +from openstef_models.utils.feature_selection import Include + + +def _dataset() -> TimeSeriesDataset: + data = pd.DataFrame( + {"load": [100.0, 200.0, 300.0], "temperature": [20.0, 25.0, 30.0]}, + index=pd.date_range("2025-01-01", periods=3, freq="h"), + ) + return TimeSeriesDataset(data, timedelta(hours=1)) + + +def test_shape_preserving_transformer_scales_in_place(): + """A scaler keeps the same columns and reports no added features.""" + dataset = _dataset() + adapter = SklearnTransformAdapter(transformer_class="sklearn.preprocessing.StandardScaler") + + adapter.fit(dataset) + result = adapter.transform(dataset) + + assert set(result.data.columns) == {"load", "temperature"} + assert result.data["load"].mean() == pytest.approx(0.0, abs=1e-9) + assert result.data["load"].std(ddof=0) == pytest.approx(1.0) + assert adapter.features_added() == [] + + +def test_shape_changing_transformer_replaces_features_with_outputs(): + """PCA replaces the input features with its components (the added features).""" + adapter = SklearnTransformAdapter( + transformer_class="sklearn.decomposition.PCA", + transformer_params={"n_components": 2, "random_state": 0}, + ) + dataset = _dataset() + + adapter.fit(dataset) + result = adapter.transform(dataset) + + added = adapter.features_added() + assert len(added) == 2 + # the two input features are gone, replaced by the two components + assert set(result.data.columns) == set(added) + assert "load" not in result.data.columns + + +def test_unselected_features_pass_through_unchanged(): + """Only the selected features are transformed; the rest pass through untouched.""" + dataset = _dataset() + adapter = SklearnTransformAdapter( + transformer_class="sklearn.preprocessing.StandardScaler", + selection=Include("load"), + ) + + adapter.fit(dataset) + result = adapter.transform(dataset) + + assert result.data["temperature"].tolist() == [20.0, 25.0, 30.0] + assert result.data["load"].mean() == pytest.approx(0.0, abs=1e-9) + + +def test_transform_before_fit_raises(): + """transform() before fit() raises NotFittedError.""" + adapter = SklearnTransformAdapter(transformer_class="sklearn.preprocessing.StandardScaler") + + with pytest.raises(NotFittedError): + adapter.transform(_dataset()) + + +def test_config_round_trips_and_rebuilds_transformer(): + """The config is serializable and reconstructs an equivalent transformer.""" + adapter = SklearnTransformAdapter( + transformer_class="sklearn.decomposition.PCA", + transformer_params={"n_components": 3}, + ) + + restored = SklearnTransformAdapter.model_validate(adapter.model_dump()) + + assert restored.transformer_class == "sklearn.decomposition.PCA" + assert restored.transformer_params == {"n_components": 3} + assert type(restored._transformer).__name__ == "PCA" + assert restored._transformer.n_components == 3 + + +def test_non_sklearn_transformer_class_raises(): + """A non-sklearn transformer_class is rejected at construction.""" + with pytest.raises(ValueError, match="scikit-learn"): + SklearnTransformAdapter(transformer_class="not_a_real_module.Nope") + + +def test_unknown_sklearn_class_raises(): + """A sklearn.* path pointing at a non-existent class is rejected.""" + with pytest.raises(ValueError, match="not found"): + SklearnTransformAdapter(transformer_class="sklearn.preprocessing.NotARealTransformer")