mirror of
https://github.com/webclinic017/drift.git
synced 2026-08-21 23:08:09 +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)
|
forward_returns = create_forward_returns(returns, self.time_horizon)
|
||||||
cutoff_point = returns.index[-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]
|
event_candidates = forward_returns[event_start_times]
|
||||||
|
|
||||||
def get_bins_threeway(x):
|
def get_bins_threeway(x):
|
||||||
|
|||||||
@@ -17,7 +17,7 @@ class FixedTimeHorionThreeClassImbalancedEventLabeller(EventLabeller):
|
|||||||
|
|
||||||
forward_returns = create_forward_returns(returns, self.time_horizon)
|
forward_returns = create_forward_returns(returns, self.time_horizon)
|
||||||
cutoff_point = returns.index[-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]
|
event_candidates = forward_returns[event_start_times]
|
||||||
|
|
||||||
def get_bins_threeway(x):
|
def get_bins_threeway(x):
|
||||||
|
|||||||
@@ -18,7 +18,7 @@ class FixedTimeHorionTwoClassEventLabeller(EventLabeller):
|
|||||||
|
|
||||||
forward_returns = create_forward_returns(returns, self.time_horizon)
|
forward_returns = create_forward_returns(returns, self.time_horizon)
|
||||||
cutoff_point = returns.index[-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]
|
event_candidates = forward_returns[event_start_times]
|
||||||
|
|
||||||
def get_class_binary(x: float) -> int:
|
def get_class_binary(x: float) -> int:
|
||||||
|
|||||||
+1
-1
@@ -27,7 +27,7 @@ def label_data(
|
|||||||
"% of overlapping events",
|
"% of overlapping events",
|
||||||
)
|
)
|
||||||
|
|
||||||
X = X.filter(items=events.index, axis=0)
|
X = X.loc[events.index]
|
||||||
y = events["label"]
|
y = events["label"]
|
||||||
forward_returns = events["returns"]
|
forward_returns = events["returns"]
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user