mirror of
https://github.com/webclinic017/drift.git
synced 2026-07-27 18:57:55 +00:00
fix(Labeling): use X.loc instead of X.filter, remove Nan from event_start_times
This commit is contained in:
@@ -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):
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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:
|
||||
|
||||
+1
-1
@@ -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"]
|
||||
|
||||
|
||||
Reference in New Issue
Block a user