SIENTIAPDE-1171
Refactor model_repository and enhance test coverage for MLFlow functionalities - Updated model_repository to ensure the 'temp' directory is created if it doesn't exist using `exist_ok=True`. - Added new tests for retraining and updating production models, including error handling scenarios. - Improved existing tests for model management workflows to ensure robustness and reliability.
This commit is contained in:
@@ -164,6 +164,18 @@ def test_get_experiment_last_run(mlflow, mlflow_repository):
|
||||
assert output == '2'
|
||||
|
||||
|
||||
@patch('laborious.utils.repository.model_repository.mlflow')
|
||||
def test_get_experiment_last_run_error(mlflow, mlflow_repository):
|
||||
mlflow.search_runs.return_value = []
|
||||
|
||||
try:
|
||||
mlflow_repository.get_experiment_last_run(0)
|
||||
except ValueError as e:
|
||||
assert str(e) == 'Runs is not a pandas DataFrame'
|
||||
else:
|
||||
assert False
|
||||
|
||||
|
||||
@patch('laborious.utils.repository.model_repository.mlflow.sklearn')
|
||||
@patch('laborious.utils.repository.model_repository.mlflow.set_experiment')
|
||||
def test_create_model_experiment(set_experiment, sklearn, mlflow_repository):
|
||||
@@ -233,14 +245,6 @@ def test_create_model_experiment(set_experiment, sklearn, mlflow_repository):
|
||||
@patch('laborious.utils.repository.model_repository.mlflow.log_artifact')
|
||||
def test_perform_model_retrain(log_artifact, log_model, log_param, start_run, mlflow_repository):
|
||||
|
||||
_os_remove = patch(
|
||||
'laborious.utils.repository.model_repository.remove')
|
||||
_os_makedirs = patch(
|
||||
'laborious.utils.repository.model_repository.makedirs')
|
||||
_os_path_exists = patch(
|
||||
'laborious.utils.repository.model_repository.path.exists',
|
||||
MagicMock(return_value=False))
|
||||
|
||||
prediction_model_mock = MagicMock()
|
||||
data_model_mock = MagicMock()
|
||||
experiment = 'test'
|
||||
@@ -277,6 +281,27 @@ def test_perform_model_retrain(log_artifact, log_model, log_param, start_run, ml
|
||||
assert output == ("Model retrained successfully", experiment)
|
||||
|
||||
|
||||
def test_retrain_model(mlflow_repository):
|
||||
data = MagicMock()
|
||||
model_name = 'test'
|
||||
|
||||
mlflow_repository.create_model_experiment = MagicMock(
|
||||
return_value=('data_model', 'prediction_model', '0'))
|
||||
|
||||
mlflow_repository.perform_model_retrain = MagicMock(
|
||||
return_value='Model retrained successfully')
|
||||
|
||||
output = mlflow_repository.retrain_model(data, model_name)
|
||||
|
||||
mlflow_repository.create_model_experiment.assert_called_once_with(
|
||||
model_name, data)
|
||||
|
||||
mlflow_repository.perform_model_retrain.assert_called_once_with(
|
||||
'data_model', 'prediction_model', '0', model_name, data)
|
||||
|
||||
assert output == 'Model retrained successfully'
|
||||
|
||||
|
||||
@patch('laborious.utils.repository.model_repository.mlflow')
|
||||
def test_update_production_model_by_run_id(mlflow, mlflow_repository):
|
||||
client_mock = MagicMock()
|
||||
@@ -312,6 +337,24 @@ def test_update_production_model_by_run_id(mlflow, mlflow_repository):
|
||||
}
|
||||
|
||||
|
||||
@patch('laborious.utils.repository.model_repository.mlflow')
|
||||
def test_update_production_model_by_run_id_error(mlflow, mlflow_repository):
|
||||
mlflow.tracking.MlflowClient.return_value = MagicMock(
|
||||
get_registered_model=MagicMock(
|
||||
return_value=MagicMock(
|
||||
latest_versions={}
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
try:
|
||||
mlflow_repository.update_production_model_by_run_id('0', 'test')
|
||||
except Exception as e:
|
||||
assert str(e) == 'Model versions is not a list'
|
||||
else:
|
||||
assert False
|
||||
|
||||
|
||||
def test_update_production_model(mlflow_repository):
|
||||
connector = mlflow_repository
|
||||
|
||||
|
||||
Reference in New Issue
Block a user