mirror of
https://github.com/webclinic017/drift.git
synced 2026-08-18 05:18:09 +00:00
refactor(Types): added nested types for Reporting (#162)
This commit is contained in:
@@ -4,11 +4,20 @@ from operator import itemgetter
|
||||
from training.primary_model import train_primary_model
|
||||
from training.meta_labeling import train_meta_labeling_model
|
||||
|
||||
from utils.encapsulation import Reporting, Asset, Single_Model, Training_Step
|
||||
from reporting.types import Reporting
|
||||
|
||||
|
||||
def primary_step(X: pd.DataFrame, y:pd.Series, original_X:pd.DataFrame, X_pca:pd.DataFrame, asset:list, target_returns:pd.Series, configs: dict, reporting: Reporting) -> tuple[Training_Step, pd.DataFrame]:
|
||||
training_step = Training_Step(level='primary')
|
||||
def primary_step(
|
||||
X: pd.DataFrame,
|
||||
y:pd.Series,
|
||||
original_X:pd.DataFrame,
|
||||
X_pca:pd.DataFrame,
|
||||
asset:list,
|
||||
target_returns:pd.Series,
|
||||
configs: dict,
|
||||
reporting: Reporting
|
||||
) -> tuple[Reporting.Training_Step, pd.DataFrame]:
|
||||
training_step = Reporting.Training_Step(level='primary')
|
||||
model_config, training_config, data_config = itemgetter('model_config', 'training_config', 'data_config')(configs)
|
||||
|
||||
# 3. Train Primary models
|
||||
@@ -60,8 +69,18 @@ def primary_step(X: pd.DataFrame, y:pd.Series, original_X:pd.DataFrame, X_pca:pd
|
||||
return training_step, current_predictions
|
||||
|
||||
|
||||
def secondary_step(X:pd.DataFrame, y:pd.Series, original_X:pd.DataFrame, X_pca:pd.DataFrame, current_predictions:pd.DataFrame, asset:list, target_returns:pd.Series, configs: dict, reporting: Reporting) -> Training_Step:
|
||||
training_step = Training_Step(level='secondary')
|
||||
def secondary_step(
|
||||
X:pd.DataFrame,
|
||||
y:pd.Series,
|
||||
original_X:pd.DataFrame,
|
||||
X_pca:pd.DataFrame,
|
||||
current_predictions:pd.DataFrame,
|
||||
asset:list,
|
||||
target_returns:pd.Series,
|
||||
configs: dict,
|
||||
reporting: Reporting
|
||||
) -> Reporting.Training_Step:
|
||||
training_step = Reporting.Training_Step(level='secondary')
|
||||
model_config, training_config, data_config = itemgetter('model_config', 'training_config', 'data_config')(configs)
|
||||
|
||||
# 5. Ensemble primary model predictions (If Ensemble model is present)
|
||||
@@ -90,7 +109,6 @@ def secondary_step(X:pd.DataFrame, y:pd.Series, original_X:pd.DataFrame, X_pca:p
|
||||
reporting.all_predictions = pd.concat([reporting.all_predictions, ensemble_predictions], axis=1)
|
||||
|
||||
|
||||
|
||||
if len(model_config['meta_labeling_models']) > 0:
|
||||
|
||||
# 3. Train a Meta-labeling model on the averaged level-1 model predictions
|
||||
|
||||
Reference in New Issue
Block a user