Skip to content

Commit 02afd31

Browse files
committed
feat: configure MVLR source-data filtering
1 parent 00f6a6f commit 02afd31

5 files changed

Lines changed: 164 additions & 24 deletions

File tree

openenergyid/mvlr/__init__.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,7 @@
77
MultiVariableRegressionInput,
88
MultiVariableRegressionResult,
99
OutlierFilteringDiagnostics,
10+
SourceDataFilteringParameters,
1011
ValidationParameters,
1112
)
1213
from .source_data_filtering import clean_regression_frame
@@ -18,6 +19,7 @@
1819
"MultiVariableRegressionInput",
1920
"MultiVariableRegressionResult",
2021
"OutlierFilteringDiagnostics",
22+
"SourceDataFilteringParameters",
2123
"ValidationParameters",
2224
"IndependentVariableResult",
2325
]

openenergyid/mvlr/main.py

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -14,7 +14,11 @@ def find_best_mvlr(
1414
best_filtering = None
1515
for granularity in data.granularities:
1616
frame = data.data_frame()
17-
frame, filtering = clean_regression_frame(frame, data.dependent_variable)
17+
frame, filtering = clean_regression_frame(
18+
frame,
19+
data.dependent_variable,
20+
data.source_data_filtering,
21+
)
1822
best_filtering = filtering
1923
frame = resample_input_data(data=frame, granularity=granularity)
2024
mvlr = MultiVariableLinearRegression(

openenergyid/mvlr/models.py

Lines changed: 55 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -34,6 +34,56 @@ class ValidationParameters(BaseModel):
3434
)
3535

3636

37+
class SourceDataFilteringParameters(BaseModel):
38+
"""Parameters for source-data filtering before regression fitting."""
39+
40+
enabled: bool = True
41+
minimum_retained_fraction: float = Field(
42+
0.50,
43+
ge=0,
44+
le=1,
45+
alias="minimumRetainedFraction",
46+
description="Minimum fraction of original observations that must remain after filtering.",
47+
)
48+
minimum_retained_rows: int = Field(
49+
30,
50+
ge=0,
51+
alias="minimumRetainedRows",
52+
description="Minimum number of observations that must remain after filtering.",
53+
)
54+
solar_reference_names: tuple[str, ...] = Field(
55+
("solarPowerGeneration", "solarRadiation"),
56+
alias="solarReferenceNames",
57+
description="Preferred source column names used as solar production/radiation references.",
58+
)
59+
positive_reference_median_fraction: float = Field(
60+
0.10,
61+
ge=0,
62+
alias="positiveReferenceMedianFraction",
63+
description="Fraction of the positive solar-reference median used as meaningful-solar threshold.",
64+
)
65+
minimum_positive_reference: float = Field(
66+
0.05,
67+
ge=0,
68+
alias="minimumPositiveReference",
69+
description="Minimum meaningful-solar threshold for the solar reference column.",
70+
)
71+
ratio_iqr_multiplier: float = Field(
72+
3.0,
73+
gt=0,
74+
alias="ratioIqrMultiplier",
75+
description="IQR multiplier used when MAD-based ratio filtering is unavailable.",
76+
)
77+
ratio_robust_z_threshold: float = Field(
78+
4.5,
79+
gt=0,
80+
alias="ratioRobustZThreshold",
81+
description="Robust z-score threshold for production-to-reference ratio outliers.",
82+
)
83+
84+
model_config = ConfigDict(populate_by_name=True)
85+
86+
3787
class IndependentVariableInput(BaseModel):
3888
"""
3989
Independent variable.
@@ -71,7 +121,11 @@ class MultiVariableRegressionInput(BaseModel):
71121
granularities: list[Granularity]
72122
allow_negative_predictions: bool = Field(alias="allowNegativePredictions", default=False)
73123
validation_parameters: ValidationParameters = Field(
74-
alias="validationParameters", default=ValidationParameters()
124+
alias="validationParameters", default_factory=ValidationParameters
125+
)
126+
source_data_filtering: SourceDataFilteringParameters = Field(
127+
alias="sourceDataFiltering",
128+
default_factory=SourceDataFilteringParameters,
75129
)
76130
single_use_exog_prefixes: list[str] | None = Field(
77131
# default=["HDD", "CDD", "FDD"],

openenergyid/mvlr/source_data_filtering.py

Lines changed: 42 additions & 21 deletions
Original file line numberDiff line numberDiff line change
@@ -3,15 +3,14 @@
33
import numpy as np
44
import pandas as pd
55

6-
from .models import OutlierFilteringDiagnostics
6+
from .models import OutlierFilteringDiagnostics, SourceDataFilteringParameters
77

8-
MINIMUM_RETAINED_FRACTION = 0.50
9-
MINIMUM_RETAINED_ROWS = 30
10-
SOLAR_REFERENCE_NAMES = ("solarPowerGeneration", "solarRadiation")
118

12-
13-
def _solar_reference_column(frame: pd.DataFrame) -> str | None:
14-
for name in SOLAR_REFERENCE_NAMES:
9+
def _solar_reference_column(
10+
frame: pd.DataFrame,
11+
parameters: SourceDataFilteringParameters,
12+
) -> str | None:
13+
for name in parameters.solar_reference_names:
1514
if name in frame.columns:
1615
return name
1716

@@ -23,24 +22,37 @@ def _solar_reference_column(frame: pd.DataFrame) -> str | None:
2322
return None
2423

2524

26-
def _is_solar_production_model(dependent_variable: str, frame: pd.DataFrame) -> bool:
25+
def _is_solar_production_model(
26+
dependent_variable: str,
27+
frame: pd.DataFrame,
28+
parameters: SourceDataFilteringParameters,
29+
) -> bool:
2730
dependent = dependent_variable.lower()
2831
if "solarphotovoltaic" in dependent:
2932
return True
3033
if "solar" in dependent and "production" in dependent:
3134
return True
32-
return "production" in dependent and _solar_reference_column(frame) is not None
35+
return "production" in dependent and _solar_reference_column(frame, parameters) is not None
3336

3437

35-
def _positive_reference_threshold(series: pd.Series) -> float:
38+
def _positive_reference_threshold(
39+
series: pd.Series,
40+
parameters: SourceDataFilteringParameters,
41+
) -> float:
3642
positive = series[series > 0]
3743
if positive.empty:
3844
return 0.0
39-
return max(float(positive.median()) * 0.10, 0.05)
45+
return max(
46+
float(positive.median()) * parameters.positive_reference_median_fraction,
47+
parameters.minimum_positive_reference,
48+
)
4049

4150

42-
def _robust_ratio_outlier_mask(ratio: pd.Series) -> pd.Series:
43-
if len(ratio) < MINIMUM_RETAINED_ROWS:
51+
def _robust_ratio_outlier_mask(
52+
ratio: pd.Series,
53+
parameters: SourceDataFilteringParameters,
54+
) -> pd.Series:
55+
if len(ratio) < parameters.minimum_retained_rows:
4456
return pd.Series(False, index=ratio.index)
4557

4658
median = float(ratio.median())
@@ -51,26 +63,35 @@ def _robust_ratio_outlier_mask(ratio: pd.Series) -> pd.Series:
5163
iqr = q3 - q1
5264
if not np.isfinite(iqr) or iqr <= 0:
5365
return pd.Series(False, index=ratio.index)
54-
return (ratio < q1 - 3.0 * iqr) | (ratio > q3 + 3.0 * iqr)
66+
return (ratio < q1 - parameters.ratio_iqr_multiplier * iqr) | (
67+
ratio > q3 + parameters.ratio_iqr_multiplier * iqr
68+
)
5569

5670
robust_z = 0.6745 * (ratio - median).abs() / mad
57-
return robust_z > 4.5
71+
return robust_z > parameters.ratio_robust_z_threshold
5872

5973

6074
def clean_regression_frame(
6175
frame: pd.DataFrame,
6276
dependent_variable: str,
77+
parameters: SourceDataFilteringParameters | None = None,
6378
) -> tuple[pd.DataFrame, OutlierFilteringDiagnostics]:
6479
"""Remove obvious bad source observations before fitting a regression model."""
6580

81+
parameters = parameters or SourceDataFilteringParameters()
6682
original_count = len(frame)
6783
diagnostics = OutlierFilteringDiagnostics(
84+
enabled=parameters.enabled,
6885
originalObservationCount=original_count,
6986
retainedObservationCount=original_count,
7087
removedObservationCount=0,
7188
applied=False,
7289
)
7390

91+
if not parameters.enabled:
92+
diagnostics.reason = "source-data filtering disabled"
93+
return frame, diagnostics
94+
7495
if original_count == 0 or dependent_variable not in frame.columns:
7596
diagnostics.reason = "empty frame or missing dependent variable"
7697
return frame, diagnostics
@@ -82,7 +103,7 @@ def clean_regression_frame(
82103
diagnostics.removed_non_finite_count = int((keep & ~finite_mask).sum())
83104
keep &= finite_mask
84105

85-
if not _is_solar_production_model(dependent_variable, numeric_frame):
106+
if not _is_solar_production_model(dependent_variable, numeric_frame, parameters):
86107
cleaned = numeric_frame.loc[keep].copy()
87108
diagnostics.retained_observation_count = len(cleaned)
88109
diagnostics.removed_observation_count = original_count - len(cleaned)
@@ -95,30 +116,30 @@ def clean_regression_frame(
95116
diagnostics.removed_negative_count = int((keep & negative_mask).sum())
96117
keep &= ~negative_mask
97118

98-
solar_column = _solar_reference_column(numeric_frame)
119+
solar_column = _solar_reference_column(numeric_frame, parameters)
99120
if solar_column is not None:
100121
solar_reference = numeric_frame[solar_column]
101-
solar_threshold = _positive_reference_threshold(solar_reference[keep])
122+
solar_threshold = _positive_reference_threshold(solar_reference[keep], parameters)
102123

103124
zero_with_solar_mask = (y <= 0) & (solar_reference > solar_threshold)
104125
diagnostics.removed_zero_with_solar_count = int((keep & zero_with_solar_mask).sum())
105126
keep &= ~zero_with_solar_mask
106127

107128
ratio_candidates = keep & (y > 0) & (solar_reference > solar_threshold)
108129
ratios = y[ratio_candidates] / solar_reference[ratio_candidates]
109-
ratio_outliers = _robust_ratio_outlier_mask(ratios)
130+
ratio_outliers = _robust_ratio_outlier_mask(ratios, parameters)
110131
diagnostics.removed_ratio_outlier_count = int(ratio_outliers.sum())
111132
keep.loc[ratio_outliers[ratio_outliers].index] = False
112133

113134
cleaned = numeric_frame.loc[keep].copy()
114135
retained_count = len(cleaned)
115136
removed_count = original_count - retained_count
116137

117-
if retained_count < MINIMUM_RETAINED_ROWS:
138+
if retained_count < parameters.minimum_retained_rows:
118139
diagnostics.reason = "too few observations retained after filtering"
119140
return numeric_frame, diagnostics
120141

121-
if retained_count / original_count < MINIMUM_RETAINED_FRACTION:
142+
if retained_count / original_count < parameters.minimum_retained_fraction:
122143
diagnostics.reason = "too much source data would be removed"
123144
return numeric_frame, diagnostics
124145

tests/mvlr/test_source_data_filtering.py

Lines changed: 60 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,7 @@
77
from openenergyid.models import TimeDataFrame
88
from openenergyid.mvlr import (
99
MultiVariableRegressionInput,
10+
SourceDataFilteringParameters,
1011
clean_regression_frame,
1112
find_best_mvlr,
1213
)
@@ -30,7 +31,9 @@ def _solar_regression_input(
3031

3132
if zero_slice:
3233
production.iloc[zero_slice] = 0.0
33-
for idx, value in (spikes or {45: 70.0, 60: 75.0, 75: 80.0}).items():
34+
if spikes is None:
35+
spikes = {45: 70.0, 60: 75.0, 75: 80.0}
36+
for idx, value in spikes.items():
3437
production.iloc[idx] = value
3538

3639
frame = pd.DataFrame(
@@ -104,6 +107,62 @@ def test_clean_regression_frame_keeps_original_data_when_filtering_too_much() ->
104107
assert len(cleaned) == 90
105108

106109

110+
def test_clean_regression_frame_uses_filtering_parameters() -> None:
111+
"""Retained-data guardrails should be caller-configurable."""
112+
data = _solar_regression_input(zero_slice=slice(0, 55), spikes={})
113+
frame = data.data_frame()
114+
115+
cleaned, diagnostics = clean_regression_frame(
116+
frame,
117+
DEPENDENT,
118+
SourceDataFilteringParameters(
119+
minimum_retained_fraction=0.30,
120+
ratio_robust_z_threshold=999.0,
121+
),
122+
)
123+
124+
assert diagnostics.applied
125+
assert diagnostics.removed_zero_with_solar_count == 55
126+
assert diagnostics.removed_observation_count == 55
127+
assert len(cleaned) == 35
128+
129+
130+
def test_source_data_filtering_parameters_support_json_aliases() -> None:
131+
"""Filtering parameters should be usable from API-shaped input."""
132+
parameters = SourceDataFilteringParameters.model_validate(
133+
{
134+
"enabled": False,
135+
"minimumRetainedRows": 12,
136+
"minimumRetainedFraction": 0.25,
137+
"solarReferenceNames": ["customSolarReference"],
138+
"ratioRobustZThreshold": 8.0,
139+
}
140+
)
141+
142+
assert not parameters.enabled
143+
assert parameters.minimum_retained_rows == 12
144+
assert parameters.minimum_retained_fraction == 0.25
145+
assert parameters.solar_reference_names == ("customSolarReference",)
146+
assert parameters.model_dump(by_alias=True)["ratioRobustZThreshold"] == 8.0
147+
148+
149+
def test_clean_regression_frame_can_be_disabled() -> None:
150+
"""Filtering can be disabled without changing the source frame."""
151+
data = _solar_regression_input()
152+
frame = data.data_frame()
153+
154+
cleaned, diagnostics = clean_regression_frame(
155+
frame,
156+
DEPENDENT,
157+
SourceDataFilteringParameters(enabled=False),
158+
)
159+
160+
assert not diagnostics.enabled
161+
assert not diagnostics.applied
162+
assert diagnostics.reason == "source-data filtering disabled"
163+
assert cleaned is frame
164+
165+
107166
def test_clean_regression_frame_removes_non_finite_rows_for_non_solar_models() -> None:
108167
"""Generic MVLR cleaning should drop non-finite observations."""
109168
index = pd.date_range("2025-04-01", periods=40, freq="D", tz="Europe/Brussels")

0 commit comments

Comments
 (0)