11import pandas as pd
22from .base import Transformer
33
4+ _UNIT_THRESHOLDS = [(86400 , "d" ), (3600 , "h" ), (60 , "min" ), (1 , "s" )]
5+
46
57class 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 )
0 commit comments