Files
drift/models/base.py
T
Mark Aron Szulyovszky 79d84cf0a3 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()
2021-12-21 10:30:09 +01:00

38 lines
731 B
Python

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))