Skip to content

Commit 32dc1ea

Browse files
committed
Dynamic label_shift dataset generation
1 parent d4df1d6 commit 32dc1ea

2 files changed

Lines changed: 359 additions & 138 deletions

File tree

Lines changed: 156 additions & 53 deletions
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,9 @@
11
from collections.abc import Iterable
2+
from pathlib import Path
23

34
import numpy as np
45
import polars as pl
6+
import polars.selectors as cs
57
import tensorflow as tf
68
from 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

147209
def 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

Comments
 (0)