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))
|
||||
Reference in New Issue
Block a user