Skip to content

Commit 39877ef

Browse files
authored
Merge pull request #21 from dawidlinek/fix/lag-transformer
Fix/lag transformer
2 parents e5ea5fe + 9aed1b6 commit 39877ef

3 files changed

Lines changed: 29 additions & 41 deletions

File tree

Lines changed: 27 additions & 40 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,8 @@
11
import pandas as pd
22
from .base import Transformer
33

4+
_UNIT_THRESHOLDS = [(86400, "d"), (3600, "h"), (60, "min"), (1, "s")]
5+
46

57
class LagTransformer(Transformer):
68
"""Transform columns by shifting them to create lagged features.
@@ -15,20 +17,15 @@ class LagTransformer(Transformer):
1517
>>> result = transformer.transform(df)
1618
"""
1719

18-
_FREQ_MAPPING = {
19-
"day": "1D",
20-
"days": "1D",
21-
"d": "1D",
22-
"hour": "1h",
23-
"hours": "1h",
24-
"h": "1h",
25-
"minute": "1min",
26-
"minutes": "1min",
27-
"min": "1min",
28-
"second": "1s",
29-
"seconds": "1s",
30-
"s": "1s",
31-
}
20+
_FREQ_MAPPING: dict[str, tuple[str, str]] = {}
21+
for _aliases, _td, _unit in [
22+
(("day", "days", "d", "1d"), "1D", "d"),
23+
(("hour", "hours", "h", "1h"), "1h", "h"),
24+
(("minute", "minutes", "min", "1min"), "1min", "min"),
25+
(("second", "seconds", "s", "1s"), "1s", "s"),
26+
]:
27+
for _alias in _aliases:
28+
_FREQ_MAPPING[_alias] = (_td, _unit)
3229

3330
def __init__(
3431
self,
@@ -39,37 +36,29 @@ def __init__(
3936
self.columns = [columns] if isinstance(columns, str) else columns
4037
self.lags = [lags] if isinstance(lags, int) else list(lags)
4138
self.freq = freq
42-
freq_normalized = self._FREQ_MAPPING.get(freq.lower(), freq)
43-
self._freq = pd.Timedelta(freq_normalized)
44-
self._validate()
4539

46-
def _validate(self) -> None:
4740
if not self.lags:
4841
raise ValueError("At least one lag value must be provided")
49-
5042
if not self.columns:
5143
raise ValueError("At least one column must be provided")
5244

53-
def _get_timedelta(self, lag: int) -> pd.Timedelta:
54-
return self._freq * lag
45+
mapping = self._FREQ_MAPPING.get(freq.lower())
46+
if mapping:
47+
freq_normalized, self._freq_unit = mapping
48+
else:
49+
freq_normalized, self._freq_unit = freq, None
50+
self._freq = pd.Timedelta(freq_normalized)
5551

5652
def _format_lag_name(self, column: str, lag: int) -> str:
57-
total_td = self._freq * abs(lag)
58-
total_seconds = int(total_td.total_seconds())
59-
60-
if total_seconds % 86400 == 0:
61-
value = total_seconds // 86400
62-
unit = "d"
63-
elif total_seconds % 3600 == 0:
64-
value = total_seconds // 3600
65-
unit = "h"
66-
elif total_seconds % 60 == 0:
67-
value = total_seconds // 60
68-
unit = "min"
53+
if self._freq_unit:
54+
unit, value = self._freq_unit, abs(lag)
6955
else:
70-
value = total_seconds
71-
unit = "s"
72-
56+
total_seconds = int((self._freq * abs(lag)).total_seconds())
57+
value, unit = next(
58+
(total_seconds // d, u)
59+
for d, u in _UNIT_THRESHOLDS
60+
if total_seconds % d == 0
61+
)
7362
sign = "-" if lag >= 0 else "+"
7463
return f"{column}_{unit}{sign}{value}"
7564

@@ -79,9 +68,7 @@ def transform(self, df: pd.DataFrame) -> pd.DataFrame:
7968
series = df[column]
8069
for lag in self.lags:
8170
name = self._format_lag_name(column, lag)
82-
timedelta = self._get_timedelta(lag)
83-
shifted_index = df.index + timedelta
84-
shifted_series = pd.Series(series.values, index=shifted_index)
85-
lagged_data[name] = shifted_series.reindex(df.index)
71+
shifted = pd.Series(series.values, index=df.index + self._freq * lag)
72+
lagged_data[name] = shifted.reindex(df.index)
8673

8774
return pd.concat([df, pd.DataFrame(lagged_data, index=df.index)], axis=1)

examples/01_data_pipeline_only.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -40,6 +40,7 @@
4040
)
4141
)
4242
.add_transformer(TimezoneTransformer(target_tz="Europe/Warsaw"))
43+
.add_transformer(ResampleTransformer(freq="1h", columns=["load_forecast_daily_min","load_forecast_daily_max"], method="ffill"))
4344
.add_transformer(ResampleTransformer(freq="1h"))
4445
.add_transformer(
4546
LagTransformer(

uv.lock

Lines changed: 1 addition & 1 deletion
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

0 commit comments

Comments
 (0)