mirror of
https://github.com/webclinic017/drift.git
synced 2026-08-15 11:58:07 +00:00
feat(Events): added EventFilter, EventLabeller (#186)
This commit is contained in:
@@ -0,0 +1,60 @@
|
||||
|
||||
|
||||
from ..types import EventFilter
|
||||
from data_loader.types import ReturnSeries
|
||||
import pandas as pd
|
||||
|
||||
class CUSUMVolatilityEventFilter(EventFilter):
|
||||
|
||||
def __init__(self, vol_period: int):
|
||||
self.vol_period = vol_period
|
||||
|
||||
def get_event_start_times(self, returns: ReturnSeries) -> pd.DatetimeIndex:
|
||||
|
||||
rolling_vol = returns.rolling(self.vol_period).std() * 0.15
|
||||
|
||||
filtered_indices = []
|
||||
pos_threshold = 0
|
||||
neg_threshold = 0
|
||||
diff = returns.diff()
|
||||
for index in diff.index[1:]:
|
||||
pos_threshold, neg_threshold = (
|
||||
max(0, pos_threshold + diff.loc[index]),
|
||||
min(0, neg_threshold + diff.loc[index]),
|
||||
)
|
||||
|
||||
if neg_threshold < -rolling_vol[index]:
|
||||
neg_threshold = 0
|
||||
filtered_indices.append(index)
|
||||
|
||||
elif pos_threshold > rolling_vol[index]:
|
||||
pos_threshold = 0
|
||||
filtered_indices.append(index)
|
||||
|
||||
return pd.DatetimeIndex(filtered_indices)
|
||||
|
||||
class CUSUMFixedEventFilter(EventFilter):
|
||||
|
||||
def __init__(self, threshold: float):
|
||||
self.threshold = threshold
|
||||
|
||||
def get_event_start_times(self, returns: ReturnSeries) -> pd.DatetimeIndex:
|
||||
filtered_indices = []
|
||||
pos_threshold = 0
|
||||
neg_threshold = 0
|
||||
diff = returns.diff()
|
||||
for index in diff.index[1:]:
|
||||
pos_threshold, neg_threshold = (
|
||||
max(0, pos_threshold + diff.loc[index]),
|
||||
min(0, neg_threshold + diff.loc[index]),
|
||||
)
|
||||
|
||||
if neg_threshold < -self.threshold:
|
||||
neg_threshold = 0
|
||||
filtered_indices.append(index)
|
||||
|
||||
elif pos_threshold > self.threshold:
|
||||
pos_threshold = 0
|
||||
filtered_indices.append(index)
|
||||
|
||||
return pd.DatetimeIndex(filtered_indices)
|
||||
@@ -0,0 +1,9 @@
|
||||
from ..types import EventFilter
|
||||
from data_loader.types import ReturnSeries
|
||||
import pandas as pd
|
||||
|
||||
class NoEventFilter(EventFilter):
|
||||
|
||||
def get_event_start_times(self, returns: ReturnSeries) -> pd.DatetimeIndex:
|
||||
return returns.index
|
||||
|
||||
@@ -0,0 +1,8 @@
|
||||
from .event_filters.nofilter import NoEventFilter
|
||||
from .event_filters.cusum import CUSUMVolatilityEventFilter, CUSUMFixedEventFilter
|
||||
|
||||
eventfilters_map = dict(
|
||||
none = NoEventFilter(),
|
||||
cusum_vol = CUSUMVolatilityEventFilter(vol_period = 20),
|
||||
cusum_fixed = CUSUMFixedEventFilter(threshold = 0.05)
|
||||
)
|
||||
@@ -0,0 +1,44 @@
|
||||
from data_loader.types import ForwardReturnSeries
|
||||
from ..types import EventLabeller, EventsDataFrame
|
||||
import pandas as pd
|
||||
|
||||
class FixedTimeHorionThreeClassBalancedEventLabeller(EventLabeller):
|
||||
|
||||
time_horizon: int
|
||||
|
||||
def __init__(self, time_horizon: int = 1):
|
||||
self.time_horizon = time_horizon
|
||||
|
||||
def label_events(self, event_start_times: pd.DatetimeIndex, forward_returns: ForwardReturnSeries) -> EventsDataFrame:
|
||||
|
||||
event_candidates = forward_returns[event_start_times]
|
||||
|
||||
def get_bins_threeway(x):
|
||||
bins = pd.qcut(event_candidates, 3, retbins=True, duplicates = 'drop')[1]
|
||||
|
||||
if len(bins) != 4:
|
||||
# if we don't have enough data for the quantiles, we'll need to add hard-coded values
|
||||
lower_bound = bins[0]
|
||||
upper_bound = bins[-1]
|
||||
bins = [lower_bound] + [-0.02, 0.02] + [upper_bound]
|
||||
return bins
|
||||
bins = get_bins_threeway(event_candidates)
|
||||
|
||||
def map_class_threeway(current_value):
|
||||
lower_threshold = bins[1]
|
||||
upper_threshold = bins[2]
|
||||
if current_value <= lower_threshold:
|
||||
return -1
|
||||
elif current_value > lower_threshold and current_value < upper_threshold:
|
||||
return 0
|
||||
else:
|
||||
return 1
|
||||
labels = event_candidates.map(map_class_threeway)
|
||||
|
||||
return pd.DataFrame({
|
||||
'start': event_start_times,
|
||||
'end': event_start_times + pd.Timedelta(days=self.time_horizon),
|
||||
'label': labels,
|
||||
'returns': forward_returns[event_start_times]
|
||||
})
|
||||
|
||||
@@ -0,0 +1,44 @@
|
||||
from ..types import EventLabeller, EventsDataFrame, ForwardReturnSeries
|
||||
import pandas as pd
|
||||
|
||||
class FixedTimeHorionThreeClassImbalancedEventLabeller(EventLabeller):
|
||||
|
||||
time_horizon: int
|
||||
|
||||
def __init__(self, time_horizon: int = 1):
|
||||
self.time_horizon = time_horizon
|
||||
|
||||
def label_events(self, event_start_times: pd.DatetimeIndex, forward_returns: ForwardReturnSeries) -> EventsDataFrame:
|
||||
|
||||
event_candidates = forward_returns[event_start_times]
|
||||
|
||||
def get_bins_threeway(x):
|
||||
bins = pd.qcut(event_candidates, 4, retbins=True, duplicates = 'drop')[1]
|
||||
|
||||
if len(bins) != 5:
|
||||
# if we don't have enough data for the quantiles, we'll need to add hard-coded values
|
||||
lower_bound = bins[0]
|
||||
upper_bound = bins[-1]
|
||||
bins = [lower_bound] + [-0.02, 0.0, 0.02] + [upper_bound]
|
||||
return bins
|
||||
bins = get_bins_threeway(event_candidates)
|
||||
|
||||
def map_class_threeway(current_value):
|
||||
lower_threshold = bins[1]
|
||||
upper_threshold = bins[3]
|
||||
if current_value <= lower_threshold:
|
||||
return -1
|
||||
elif current_value > lower_threshold and current_value < upper_threshold:
|
||||
return 0
|
||||
else:
|
||||
return 1
|
||||
labels = event_candidates.map(map_class_threeway)
|
||||
|
||||
return pd.DataFrame({
|
||||
'start': event_start_times,
|
||||
'end': event_start_times + pd.Timedelta(days=self.time_horizon),
|
||||
'label': labels,
|
||||
'returns': forward_returns[event_start_times]
|
||||
})
|
||||
|
||||
|
||||
@@ -0,0 +1,29 @@
|
||||
|
||||
|
||||
from ..types import EventLabeller, EventsDataFrame, ForwardReturnSeries
|
||||
import pandas as pd
|
||||
|
||||
class FixedTimeHorionTwoClassEventLabeller(EventLabeller):
|
||||
|
||||
time_horizon: int
|
||||
|
||||
def __init__(self, time_horizon: int = 1):
|
||||
self.time_horizon = time_horizon
|
||||
|
||||
def label_events(self, event_start_times: pd.DatetimeIndex, forward_returns: ForwardReturnSeries) -> EventsDataFrame:
|
||||
|
||||
event_candidates = forward_returns[event_start_times]
|
||||
|
||||
def get_class_binary(x: float) -> int:
|
||||
return -1 if x <= 0.0 else 1
|
||||
labels = event_candidates.map(get_class_binary)
|
||||
|
||||
return pd.DataFrame({
|
||||
'start': event_start_times,
|
||||
'end': event_start_times + pd.Timedelta(days=self.time_horizon),
|
||||
'label': labels,
|
||||
'returns': forward_returns[event_start_times]
|
||||
})
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,9 @@
|
||||
from .labellers.fixed_time_three_class_balanced import FixedTimeHorionThreeClassBalancedEventLabeller
|
||||
from .labellers.fixed_time_three_class_imbalanced import FixedTimeHorionThreeClassImbalancedEventLabeller
|
||||
from .labellers.fixed_time_two_class import FixedTimeHorionTwoClassEventLabeller
|
||||
|
||||
labellers_map = dict(
|
||||
two_class = FixedTimeHorionTwoClassEventLabeller(),
|
||||
three_class_balanced = FixedTimeHorionThreeClassBalancedEventLabeller(),
|
||||
three_class_imbalanced = FixedTimeHorionThreeClassImbalancedEventLabeller()
|
||||
)
|
||||
@@ -0,0 +1,17 @@
|
||||
from .types import EventFilter, EventLabeller, EventsDataFrame
|
||||
from data_loader.types import ForwardReturnSeries, XDataFrame, ReturnSeries, ySeries
|
||||
|
||||
def label_data(
|
||||
event_filter: EventFilter,
|
||||
event_labeller: EventLabeller,
|
||||
X: XDataFrame,
|
||||
returns: ReturnSeries,
|
||||
forward_returns: ForwardReturnSeries) -> tuple[EventsDataFrame, XDataFrame, ySeries, ForwardReturnSeries]:
|
||||
event_start_times = event_filter.get_event_start_times(returns)
|
||||
events = event_labeller.label_events(event_start_times, forward_returns)
|
||||
|
||||
X = X.filter(items = events.index, axis = 0)
|
||||
y = events['label']
|
||||
forward_returns = events['returns']
|
||||
|
||||
return events, X, y, forward_returns
|
||||
@@ -0,0 +1,28 @@
|
||||
from data_loader.types import ReturnSeries, ForwardReturnSeries
|
||||
from abc import ABC, abstractmethod
|
||||
import pandas as pd
|
||||
import pandera as pa
|
||||
from pandera.typing import DataFrame, Series
|
||||
|
||||
class EventFilter(ABC):
|
||||
|
||||
@abstractmethod
|
||||
def get_event_start_times(self, returns: ReturnSeries) -> pd.DatetimeIndex:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class EventSchema(pa.SchemaModel):
|
||||
start: Series[pd.Timestamp]
|
||||
end: Series[pd.Timestamp]
|
||||
label: Series[int]
|
||||
returns: Series[float]
|
||||
|
||||
EventsDataFrame = DataFrame[EventSchema]
|
||||
|
||||
|
||||
class EventLabeller(ABC):
|
||||
|
||||
@abstractmethod
|
||||
def label_events(self, event_start_times: pd.DatetimeIndex, forward_returns: ForwardReturnSeries) -> EventsDataFrame:
|
||||
raise NotImplementedError
|
||||
|
||||
Reference in New Issue
Block a user