feat: add experiment_name to TrainModelResult and update training logic
- Introduced experiment_name parameter in TrainModelResult to enhance tracking of training experiments. - Updated the Training class to utilize run_name and experiment_name for improved MLflow run management.
This commit is contained in:
@@ -283,17 +283,17 @@ class Training(SientiaMonitoring):
|
|||||||
self.info(f'Starting MLflow run for {train_params.model_type}', metadata)
|
self.info(f'Starting MLflow run for {train_params.model_type}', metadata)
|
||||||
with self.mlflow_repository.start_run(
|
with self.mlflow_repository.start_run(
|
||||||
model_name=train_params.model_name,
|
model_name=train_params.model_name,
|
||||||
run_name=None,
|
run_name=train_result.run_name,
|
||||||
experiment_name=f'{train_params.model_name}_experiment',
|
experiment_name=train_result.experiment_name,
|
||||||
tags=None,
|
tags=None,
|
||||||
metadata=metadata,
|
metadata=metadata,
|
||||||
) as run_info:
|
) as run_info:
|
||||||
train_result.run_name = run_info.run_name
|
|
||||||
train_result.run_id = run_info.run_id
|
train_result.run_id = run_info.run_id
|
||||||
self._persist_training_artifacts(train_result, train_params, wrapper, metadata)
|
self._persist_training_artifacts(train_result, train_params, wrapper, metadata)
|
||||||
|
|
||||||
return {
|
return {
|
||||||
'run_name': train_result.run_name,
|
'run_name': train_result.run_name,
|
||||||
|
'experiment_name': train_result.experiment_name,
|
||||||
'run_id': train_result.run_id,
|
'run_id': train_result.run_id,
|
||||||
'run_dir': train_result.run_dir,
|
'run_dir': train_result.run_dir,
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -44,6 +44,7 @@ class TrainModelResult:
|
|||||||
equation: dict | None = None
|
equation: dict | None = None
|
||||||
equation_path: str | None = None
|
equation_path: str | None = None
|
||||||
run_name: str | None = None
|
run_name: str | None = None
|
||||||
|
experiment_name: str | None = None
|
||||||
run_id: str | None = None
|
run_id: str | None = None
|
||||||
report_path: str | None = None
|
report_path: str | None = None
|
||||||
train_data_path: str | None = None
|
train_data_path: str | None = None
|
||||||
|
|||||||
@@ -20,6 +20,7 @@ from os import makedirs, path
|
|||||||
from shutil import rmtree
|
from shutil import rmtree
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
|
from mlflow.entities import experiment
|
||||||
import numpy as np
|
import numpy as np
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
from sientia_do.observability.logger import Logger
|
from sientia_do.observability.logger import Logger
|
||||||
@@ -200,9 +201,10 @@ class DataManagerRepository(SientiaMonitoring):
|
|||||||
metadata,
|
metadata,
|
||||||
)
|
)
|
||||||
|
|
||||||
run_name = f'train_model_{params.model_type}_{params.model_name}_{params.experiment_run_id}'
|
experiment_name = f'train_model_{params.model_type}_{params.model_name}_{params.experiment_run_id}'
|
||||||
|
run_name = f'{experiment_name}_{datetime.now().strftime("%Y%m%d_%H%M%S")}'
|
||||||
|
|
||||||
return TrainModelResult(params=params, train_data=train_data, val_data=val_data, run_name=run_name)
|
return TrainModelResult(params=params, train_data=train_data, val_data=val_data, run_name=run_name, experiment_name=experiment_name)
|
||||||
|
|
||||||
def _as_series(self, pred: pd.DataFrame | pd.Series) -> pd.Series:
|
def _as_series(self, pred: pd.DataFrame | pd.Series) -> pd.Series:
|
||||||
if isinstance(pred, pd.Series):
|
if isinstance(pred, pd.Series):
|
||||||
|
|||||||
Reference in New Issue
Block a user