From c1d79ac8c2c28507f1dd4493eec1d0a193cde243 Mon Sep 17 00:00:00 2001 From: Mark Aron Szulyovszky Date: Sat, 12 Mar 2022 13:27:44 +0100 Subject: [PATCH] fix(Labeling): use X.loc instead of X.filter, remove Nan from event_start_times --- labeling/labellers/fixed_time_three_class_balanced.py | 2 +- labeling/labellers/fixed_time_three_class_imbalanced.py | 2 +- labeling/labellers/fixed_time_two_class.py | 2 +- labeling/process.py | 2 +- 4 files changed, 4 insertions(+), 4 deletions(-) diff --git a/labeling/labellers/fixed_time_three_class_balanced.py b/labeling/labellers/fixed_time_three_class_balanced.py index bece946..07ae563 100644 --- a/labeling/labellers/fixed_time_three_class_balanced.py +++ b/labeling/labellers/fixed_time_three_class_balanced.py @@ -18,7 +18,7 @@ class FixedTimeHorionThreeClassBalancedEventLabeller(EventLabeller): forward_returns = create_forward_returns(returns, self.time_horizon) cutoff_point = returns.index[-self.time_horizon] - event_start_times[event_start_times < cutoff_point] + event_start_times = event_start_times[event_start_times < cutoff_point] event_candidates = forward_returns[event_start_times] def get_bins_threeway(x): diff --git a/labeling/labellers/fixed_time_three_class_imbalanced.py b/labeling/labellers/fixed_time_three_class_imbalanced.py index b2278e6..d9aa107 100644 --- a/labeling/labellers/fixed_time_three_class_imbalanced.py +++ b/labeling/labellers/fixed_time_three_class_imbalanced.py @@ -17,7 +17,7 @@ class FixedTimeHorionThreeClassImbalancedEventLabeller(EventLabeller): forward_returns = create_forward_returns(returns, self.time_horizon) cutoff_point = returns.index[-self.time_horizon] - event_start_times[event_start_times < cutoff_point] + event_start_times = event_start_times[event_start_times < cutoff_point] event_candidates = forward_returns[event_start_times] def get_bins_threeway(x): diff --git a/labeling/labellers/fixed_time_two_class.py b/labeling/labellers/fixed_time_two_class.py index ca9cb17..aee5d52 100644 --- a/labeling/labellers/fixed_time_two_class.py +++ b/labeling/labellers/fixed_time_two_class.py @@ -18,7 +18,7 @@ class FixedTimeHorionTwoClassEventLabeller(EventLabeller): forward_returns = create_forward_returns(returns, self.time_horizon) cutoff_point = returns.index[-self.time_horizon] - event_start_times[event_start_times < cutoff_point] + event_start_times = event_start_times[event_start_times < cutoff_point] event_candidates = forward_returns[event_start_times] def get_class_binary(x: float) -> int: diff --git a/labeling/process.py b/labeling/process.py index d979b25..5f573e4 100644 --- a/labeling/process.py +++ b/labeling/process.py @@ -27,7 +27,7 @@ def label_data( "% of overlapping events", ) - X = X.filter(items=events.index, axis=0) + X = X.loc[events.index] y = events["label"] forward_returns = events["returns"]