mirror of
https://github.com/webclinic017/drift.git
synced 2026-08-21 23:08:09 +00:00
Feature(Speed): Python launches faster by conditionally importing models. (#169)
* feat: Added optional import of models. * fix: Models weren't wrapped into abstract class, fixed it. * chore: Deleted leftover comments. * fix: Same merge commit as on remote. * fix: System wasn't putting in RF because there was no differentiation between RF as regressor and RF as classificator. * fix(Models): use the XGBoostModel wrapper Co-authored-by: Daniel Szemerey <szemereydaniel@gmail.com> Co-authored-by: Mark Aron Szulyovszky <mark.szulyovszky@gmail.com>
This commit is contained in:
co-authored by
Daniel Szemerey
Mark Aron Szulyovszky
parent
797d45d036
commit
31dc847be1
+2
-2
@@ -72,8 +72,8 @@ def get_default_ensemble_config() -> tuple[dict, dict, dict]:
|
||||
narrow_format = False,
|
||||
)
|
||||
|
||||
regression_models = ["Lasso", "KNN", "RF"]
|
||||
classification_models = ["LR_two_class", "LDA", "NB", "RF", "XGB_two_class", "LGBM", "StaticMom"]
|
||||
regression_models = ["Lasso", "KNN", "RFR"]
|
||||
classification_models = ["LR_two_class", "LDA", "NB", "RFC", "XGB_two_class", "LGBM", "StaticMom"]
|
||||
meta_labeling_models = ['LR_two_class', 'LGBM']
|
||||
ensemble_model = 'Average'
|
||||
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
|
||||
from utils.helpers import flatten
|
||||
from feature_extractors.feature_extractor_presets import presets as feature_extractor_presets
|
||||
from models.model_map import model_map
|
||||
from models.model_map import get_model_map
|
||||
from data_loader.collections import data_collections
|
||||
|
||||
def preprocess_config(model_config:dict, training_config:dict, data_config:dict) -> tuple[dict, dict, dict]:
|
||||
@@ -21,6 +21,7 @@ def __preprocess_feature_extractors_config(data_dict: dict) -> dict:
|
||||
return data_dict
|
||||
|
||||
def __preprocess_model_config(model_config:dict, method:str) -> dict:
|
||||
model_map, _, _, _, _ = get_model_map(model_config)
|
||||
model_config['primary_models'] = [(model_name, model_map[method + '_models'][model_name]) for model_name in model_config['primary_models']]
|
||||
if len(model_config['meta_labeling_models']) > 0:
|
||||
model_config['meta_labeling_models'] = [(model_name, model_map[method + '_models'][model_name]) for model_name in model_config['meta_labeling_models']]
|
||||
|
||||
Reference in New Issue
Block a user