Skip to content
Open
Show file tree
Hide file tree
Changes from 1 commit
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand All @@ -34,4 +35,5 @@
"Scaler",
"Selector",
"Shifter",
"SklearnTransformAdapter",
]
Original file line number Diff line number Diff line change
@@ -0,0 +1,118 @@
# SPDX-FileCopyrightText: 2025 Contributors to the OpenSTEF project <openstef@lfenergy.org>
#
# 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 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:
module_path, _, class_name = self.transformer_class.rpartition(".")
if not module_path:
msg = f"transformer_class must be a fully qualified import path, got {self.transformer_class!r}."
raise ValueError(msg)
transformer_cls = getattr(importlib.import_module(module_path), class_name)
self._transformer = transformer_cls(**self.transformer_params)
Comment thread
Valyrian-Code marked this conversation as resolved.
if not hasattr(self._transformer, "get_feature_names_out"):
msg = f"{self.transformer_class} does not implement get_feature_names_out and cannot be adapted."
raise TypeError(msg)

@override
def fit(self, data: TimeSeriesDataset) -> None:
features = self.selection.resolve(data.feature_names)
self._transformer.fit(data.data[features])
output_names = list(self._transformer.get_feature_names_out(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 = list(self._transformer.get_feature_names_out(features))
Comment thread
Valyrian-Code marked this conversation as resolved.
Outdated
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"]
Original file line number Diff line number Diff line change
@@ -0,0 +1,97 @@
# SPDX-FileCopyrightText: 2025 Contributors to the OpenSTEF project <openstef@lfenergy.org>
#
# 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_invalid_transformer_class_raises():
"""An unimportable transformer_class fails fast at construction."""
with pytest.raises(ModuleNotFoundError):
SklearnTransformAdapter(transformer_class="not_a_real_module.Nope")
Comment thread
Valyrian-Code marked this conversation as resolved.
Outdated
Loading