mirror of
https://github.com/webclinic017/drift.git
synced 2026-08-14 03:18:06 +00:00
feat(Model): added own Model class, SkLearnModel wrapper and StaticMomentumModel (#61)
* feat(Model): added own `Model` class, SkLearnModel wrapper and StaticMomentumModel * fix(Tests): added missing Model variable * fix(Tests): added missing clone method()
This commit is contained in:
@@ -0,0 +1,38 @@
|
||||
|
||||
from typing import Literal, Optional
|
||||
from sklearn.base import clone
|
||||
|
||||
class Model:
|
||||
|
||||
# data_format: Literal['dataframe', 'numpy']
|
||||
data_scaling: Literal["scaled", "unscaled"]
|
||||
# data_format: Literal["wide", "narrow"]
|
||||
only_column: Optional[str]
|
||||
|
||||
def fit(self, X, y):
|
||||
pass
|
||||
|
||||
def predict(self, X):
|
||||
pass
|
||||
|
||||
def clone(self):
|
||||
pass
|
||||
|
||||
|
||||
class SKLearnModel(Model):
|
||||
|
||||
# data_format = 'numpy'
|
||||
data_scaling = 'scaled'
|
||||
only_column = None
|
||||
|
||||
def __init__(self, model):
|
||||
self.model = model
|
||||
|
||||
def fit(self, X, y):
|
||||
self.model.fit(X, y)
|
||||
|
||||
def predict(self, X):
|
||||
return self.model.predict(X)
|
||||
|
||||
def clone(self):
|
||||
return SKLearnModel(clone(self.model))
|
||||
@@ -0,0 +1,27 @@
|
||||
from models.base import Model
|
||||
import numpy as np
|
||||
|
||||
class StaticMomentumModel(Model):
|
||||
'''
|
||||
Model that uses only one feature: momentum. It's positive if momentum is greater than 0, otherwise it's negative.
|
||||
'''
|
||||
|
||||
data_format = 'dataframe'
|
||||
data_scaling = 'unscaled'
|
||||
only_column = 'mom'
|
||||
|
||||
def __init__(self, allow_short: bool) -> None:
|
||||
super().__init__()
|
||||
self.allow_short = allow_short
|
||||
|
||||
def fit(self, X, y):
|
||||
# This is a static model, it can' learn anything
|
||||
pass
|
||||
|
||||
def predict(self, X):
|
||||
negative_class = -1.0 if self.allow_short == True else 0.0
|
||||
prediction = 1.0 if X[-1][0] > 0 else negative_class
|
||||
return np.array([prediction])
|
||||
|
||||
def clone(self):
|
||||
return self
|
||||
Reference in New Issue
Block a user