11from collections .abc import Iterable
2+ from pathlib import Path
23
34import numpy as np
45import polars as pl
6+ import polars .selectors as cs
57import tensorflow as tf
68from polars ._typing import IntoExpr
79
@@ -14,14 +16,14 @@ class FallenDataset:
1416 input_data : pl .DataFrame
1517 labels : pl .DataFrame
1618 groups : pl .DataFrame
17- features : Iterable [IntoExpr ] | IntoExpr
19+ features : Iterable [IntoExpr ]
1820
1921 def __init__ (
2022 self ,
2123 dataframe : pl .DataFrame ,
2224 * ,
23- group_keys : list [str ],
24- features : Iterable [IntoExpr ] | IntoExpr ,
25+ group_keys : list [pl . Expr ],
26+ features : Iterable [IntoExpr ],
2527 ) -> None :
2628 self .dataframe = dataframe .drop_nulls ()
2729
@@ -70,43 +72,58 @@ def to_windowed(
7072 ],
7173 how = "horizontal" ,
7274 )
73- shifted_labels = self . dataframe . select (
74- pl . col ( "labels" ). shift ( - label_shift ). over ( "group" )
75- )[:: samples_between_windows ]
76-
77- predecessors_of_shifted_labels = self . dataframe . select (
78- pl . col ( "labels" )
79- . shift ( - ( label_shift - 1 ))
80- . over ( "group" )
81- . alias ( "label_predecessor" )
82- )[:: samples_between_windows ]
83-
84- windowed_dataframe_with_predecessors = (
75+ shifted_labels = pl . concat (
76+ [
77+ self . dataframe . select (
78+ generate_shifts (
79+ pl . col ( "labels" ), label_shift , "group" , "labels"
80+ ),
81+ )[:: samples_between_windows ]
82+ ],
83+ how = "horizontal" ,
84+ )
85+
86+ windowed_filtered_dataframe = (
8587 windowed_features .hstack (shifted_labels )
86- .hstack (predecessors_of_shifted_labels )
8788 .drop_nulls ()
89+ .sample (fraction = 1 , shuffle = True , seed = 1 )
8890 )
8991
90- windowed_dataframe = windowed_dataframe_with_predecessors .filter (
91- pl .col ("label_predecessor" ) == Label .Stable
92- ).drop ("label_predecessor" )
92+ train_test_ratio = 0.8
93+ number_of_windows = len (windowed_filtered_dataframe )
94+ split_index = int (number_of_windows * train_test_ratio )
95+ self .train_windowed_filtered_dataframe = windowed_filtered_dataframe [
96+ :split_index :
97+ ]
9398
94- n_minority = (
95- windowed_dataframe .get_column ("labels" )
96- .value_counts ()
97- .min ()
98- .select (pl .col ("count" ))
99- .item ()
99+ self .test_windowed_filtered_dataframe = windowed_filtered_dataframe [
100+ split_index + 1 :
101+ ]
102+
103+ windowed_dataframe = filter_with_labels_predecessor (
104+ self .train_windowed_filtered_dataframe , label_shift
105+ )
106+
107+ balanced_dataframe = do_class_balancing (
108+ windowed_dataframe , label_shift - 1
109+ )
110+
111+ self .input_data = balanced_dataframe .select (
112+ generate_selector_up_to_index (self .features , label_shift )
100113 )
101- balanced_df = windowed_dataframe .group_by (
102- "labels" , maintain_order = True
103- ).map_groups (lambda group : group .sample (n = n_minority , seed = 1 ))
104- shuffeled_balanced_df = balanced_df .sample (
105- fraction = 1 , shuffle = True , seed = 1
114+
115+ self .labels = balanced_dataframe .select (
116+ cs .starts_with ("labels" ) & cs .ends_with ("_" + str (label_shift ))
106117 )
107118
108- self .input_data = shuffeled_balanced_df .drop ("labels" )
109- self .labels = shuffeled_balanced_df .select (pl .col ("labels" ))
119+ cache_path = Path ("./.cache" )
120+ cache_path .mkdir (exist_ok = True )
121+ self .train_windowed_filtered_dataframe .write_parquet (
122+ cache_path .joinpath ("/train_windowed_filtered_dataframe.parquet" )
123+ )
124+ self .test_windowed_filtered_dataframe .write_parquet (
125+ cache_path .joinpath ("/test_windowed_filtered_dataframe.parquet" )
126+ )
110127
111128 def __len__ (self ) -> int :
112129 return self .groups .unique ().numel ()
@@ -121,30 +138,116 @@ def __getitem__(self, index: int) -> tuple[tf.Tensor, tf.Tensor]:
121138 mask = self .groups == index
122139 return self .input_data [mask ], self .labels [mask ]
123140
124- def get_input_tensor (self ) -> tf .Tensor :
125- num_windows = len (self .input_data )
126- window_length = self .samples_per_window
127- num_features = len (self .features )
128- return tf .convert_to_tensor (
129- np .stack ([self .input_data .to_numpy ()]).reshape (
130- (num_windows , window_length , num_features )
131- ),
132- dtype = tf .float32 ,
133- )
134-
135- def get_labels_tensor (self ) -> tf .Tensor :
136- return tf .convert_to_tensor (self .labels .to_numpy (), dtype = tf .float32 )
137141
138- def get_windows_input_tensor (self ) -> list [tf .Tensor ]:
139- windowed_input_tensor = []
140- for input_window in self .input_data .iter ():
141- windowed_input_tensor .append (
142- tf .convert_to_tensor (input_window .to_pandas ())
143- )
144- return windowed_input_tensor
142+ def get_input_tensor (
143+ input_data : pl .DataFrame , samples_per_window : int , features : list [pl .Expr ]
144+ ) -> tf .Tensor :
145+ num_windows = len (input_data )
146+ window_length = samples_per_window
147+ num_features = len (features )
148+ return tf .convert_to_tensor (
149+ np .stack ([input_data .to_numpy ()]).reshape (
150+ (num_windows , window_length , num_features )
151+ ),
152+ dtype = tf .float32 ,
153+ )
154+
155+
156+ def get_input_tensor_up_to_shift (
157+ windowed_filtered_dataframe : pl .DataFrame ,
158+ features : list [pl .Expr ],
159+ samples_per_window : int ,
160+ label_shift : int ,
161+ ) -> tf .Tensor :
162+ filtered_windowed_dataframe = filter_with_labels_predecessor (
163+ windowed_filtered_dataframe , label_shift
164+ )
165+
166+ shuffled_balanced_df = do_class_balancing (
167+ filtered_windowed_dataframe , label_shift
168+ )
169+
170+ input_data = shuffled_balanced_df .select (
171+ generate_selector_up_to_index (features , samples_per_window )
172+ )
173+
174+ num_windows = len (input_data )
175+ window_length = samples_per_window
176+ num_features = len (features )
177+ return tf .convert_to_tensor (
178+ np .stack ([input_data .to_numpy ()]).reshape (
179+ (num_windows , window_length , num_features )
180+ ),
181+ dtype = tf .float32 ,
182+ )
183+
184+
185+ def get_labels_tensor (labels : pl .DataFrame ) -> tf .Tensor :
186+ return tf .convert_to_tensor (labels .to_numpy (), dtype = tf .float32 )
187+
188+
189+ def get_labels_tensor_up_to_shift (
190+ windowed_filtered_dataframe : pl .DataFrame ,
191+ label_shift : int ,
192+ ) -> tf .Tensor :
193+ filtered_windowed_dataframe = filter_with_labels_predecessor (
194+ windowed_filtered_dataframe , label_shift
195+ )
196+
197+ shuffled_balanced_df = do_class_balancing (
198+ filtered_windowed_dataframe , label_shift
199+ )
200+
201+ labels_at_shift_index = shuffled_balanced_df .select (
202+ cs .starts_with ("labels" ) & cs .ends_with ("_" + str (label_shift ))
203+ )
204+ return tf .convert_to_tensor (
205+ labels_at_shift_index .to_numpy (), dtype = tf .float32
206+ )
145207
146208
147209def generate_lags (feature : pl .Expr , lags : int , group : str ) -> list [pl .Expr ]:
148210 return [
149- feature .shift (i ).over (group ).name .suffix (str (i )) for i in range (lags )
211+ feature .shift (i ).over (group ).name .suffix ("_" + str (i ))
212+ for i in range (lags )
213+ ]
214+
215+
216+ def generate_shifts (
217+ feature : pl .Expr , shifts : int , group : str , alias : str
218+ ) -> list [pl .Expr ]:
219+ return [
220+ feature .shift (- i ).over (group ).alias (alias + "_" + str (i ))
221+ for i in range (shifts )
222+ ]
223+
224+
225+ def filter_with_labels_predecessor (
226+ windowed_dataframe : pl .DataFrame , label_shift : int
227+ ) -> pl .DataFrame :
228+ return windowed_dataframe .filter (
229+ pl .col ("labels" + "_" + str (label_shift - 1 )) == Label .Stable
230+ )
231+
232+
233+ def do_class_balancing (
234+ filtered_windowed_dataframe : pl .DataFrame , index : int
235+ ) -> pl .DataFrame :
236+ n_minority = (
237+ filtered_windowed_dataframe .get_column ("labels" + "_" + str (index ))
238+ .value_counts ()
239+ .min ()
240+ .select (pl .col ("count" ))
241+ .item ()
242+ )
243+ return filtered_windowed_dataframe .group_by (
244+ "labels" + "_" + str (index ), maintain_order = True
245+ ).map_groups (lambda group : group .sample (n = n_minority , seed = 1 ))
246+
247+
248+ def generate_selector_up_to_index (features : list [pl .Expr ], index : int ) -> list :
249+ return [
250+ cs .starts_with (feature .meta .output_name ()) & cs .ends_with ("_" + str (i ))
251+ for feature in features
252+ for i in range (index )
150253 ]
0 commit comments