SIENTIAPDE-1430: Refactor DataPreprocessor date filtering, make rce_train radius optional, and introduce constants for model names.
This commit is contained in:
@@ -442,6 +442,36 @@ class DataPreprocessor(BaseEstimator, TransformerMixin):
|
||||
input_data = treat_nan(input_data, treatment)
|
||||
return input_data
|
||||
|
||||
def _parse_datetime(self, date_str: str | None) -> pd.Timestamp | None:
|
||||
"""Parse a date string to Timestamp, returning None on failure."""
|
||||
if not date_str:
|
||||
return None
|
||||
try:
|
||||
return pd.to_datetime(date_str)
|
||||
except (ValueError, TypeError):
|
||||
return None
|
||||
|
||||
def _filter_by_date_range(
|
||||
self, input_data: pd.DataFrame, start: pd.Timestamp | None, end: pd.Timestamp | None
|
||||
) -> pd.DataFrame:
|
||||
"""Filter DataFrame by start and end dates."""
|
||||
if start is not None:
|
||||
input_data = input_data[input_data.index >= start]
|
||||
if end is not None:
|
||||
input_data = input_data[input_data.index <= end]
|
||||
return input_data
|
||||
|
||||
def _remove_interval(self, input_data: pd.DataFrame, interval: tuple | list) -> pd.DataFrame:
|
||||
"""Remove a single interval from the DataFrame."""
|
||||
if len(interval) < 2:
|
||||
return input_data
|
||||
interval_start = self._parse_datetime(interval[0])
|
||||
interval_end = self._parse_datetime(interval[1])
|
||||
if interval_start is None or interval_end is None:
|
||||
return input_data
|
||||
mask = ~((input_data.index >= interval_start) & (input_data.index <= interval_end))
|
||||
return input_data[mask]
|
||||
|
||||
def range_selection(self, input_data: pd.DataFrame) -> pd.DataFrame:
|
||||
"""
|
||||
Filter data by date range and remove specified intervals.
|
||||
@@ -452,35 +482,13 @@ class DataPreprocessor(BaseEstimator, TransformerMixin):
|
||||
Returns:
|
||||
pandas.DataFrame: The filtered data
|
||||
"""
|
||||
# Filter by start_date and end_date
|
||||
if self.start_date:
|
||||
try:
|
||||
start = pd.to_datetime(self.start_date)
|
||||
input_data = input_data[input_data.index >= start]
|
||||
except (ValueError, TypeError):
|
||||
pass # Invalid date format, skip filtering
|
||||
start = self._parse_datetime(self.start_date)
|
||||
end = self._parse_datetime(self.end_date)
|
||||
input_data = self._filter_by_date_range(input_data, start, end)
|
||||
|
||||
if self.end_date:
|
||||
try:
|
||||
end = pd.to_datetime(self.end_date)
|
||||
input_data = input_data[input_data.index <= end]
|
||||
except (ValueError, TypeError):
|
||||
pass # Invalid date format, skip filtering
|
||||
|
||||
# Remove specified intervals
|
||||
if self.removed_intervals:
|
||||
for interval in self.removed_intervals:
|
||||
if len(interval) >= 2:
|
||||
try:
|
||||
interval_start = pd.to_datetime(interval[0])
|
||||
interval_end = pd.to_datetime(interval[1])
|
||||
mask = ~(
|
||||
(input_data.index >= interval_start)
|
||||
& (input_data.index <= interval_end)
|
||||
)
|
||||
input_data = input_data[mask]
|
||||
except (ValueError, TypeError):
|
||||
pass # Invalid date format, skip this interval
|
||||
input_data = self._remove_interval(input_data, interval)
|
||||
|
||||
return input_data
|
||||
|
||||
|
||||
Reference in New Issue
Block a user