33import numpy as np
44import 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
6074def 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
0 commit comments