SIENTIAPDE-1243: Fix: Corrected dataframe inference and mlflow repository tests, and updated activity method calls in format and export prediction tests.
This commit is contained in:
@@ -54,8 +54,8 @@ def nan_values_filter(predictions: DataFrame, _config: dict) -> bool:
|
|||||||
"""
|
"""
|
||||||
data = (
|
data = (
|
||||||
predictions.replace({None: np.nan})
|
predictions.replace({None: np.nan})
|
||||||
|
.infer_objects(copy=False)
|
||||||
.drop(columns=['timestamp'], errors='ignore')
|
.drop(columns=['timestamp'], errors='ignore')
|
||||||
.infer_objects()
|
|
||||||
)
|
)
|
||||||
|
|
||||||
if data.isna().all().all():
|
if data.isna().all().all():
|
||||||
|
|||||||
@@ -194,7 +194,7 @@ def test_get_experiment_success(mlflow, mlflow_repository):
|
|||||||
|
|
||||||
output = mlflow_repository.get_experiment('test')
|
output = mlflow_repository.get_experiment('test')
|
||||||
|
|
||||||
assert output == 0
|
assert output == '0'
|
||||||
|
|
||||||
|
|
||||||
@patch('model_manager.utils.repository.model_repository.mlflow')
|
@patch('model_manager.utils.repository.model_repository.mlflow')
|
||||||
@@ -245,7 +245,7 @@ def test_get_experiment_last_run_error(mlflow, mlflow_repository):
|
|||||||
@patch('model_manager.utils.repository.model_repository.mlflow.sklearn')
|
@patch('model_manager.utils.repository.model_repository.mlflow.sklearn')
|
||||||
@patch('model_manager.utils.repository.model_repository.mlflow.set_experiment')
|
@patch('model_manager.utils.repository.model_repository.mlflow.set_experiment')
|
||||||
def test_create_model_experiment(set_experiment, sklearn, mlflow_repository):
|
def test_create_model_experiment(set_experiment, sklearn, mlflow_repository):
|
||||||
mlflow_repository.model_serving.get_model_run_id = MagicMock(return_value='0')
|
mlflow_repository.model_serving.get_model_info = MagicMock(return_value='0')
|
||||||
mlflow_repository.model_serving.get_model_uri = MagicMock(return_value='test')
|
mlflow_repository.model_serving.get_model_uri = MagicMock(return_value='test')
|
||||||
mlflow_repository.get_experiment_by_run_id = MagicMock()
|
mlflow_repository.get_experiment_by_run_id = MagicMock()
|
||||||
|
|
||||||
@@ -268,9 +268,7 @@ def test_create_model_experiment(set_experiment, sklearn, mlflow_repository):
|
|||||||
|
|
||||||
output = mlflow_repository.create_model_experiment('test', data)
|
output = mlflow_repository.create_model_experiment('test', data)
|
||||||
|
|
||||||
mlflow_repository.model_serving.get_model_run_id.assert_called_once_with(
|
mlflow_repository.model_serving.get_model_info.assert_called_once_with('test')
|
||||||
'test', stage='Production'
|
|
||||||
)
|
|
||||||
mlflow_repository.model_serving.get_model_uri.assert_called_once_with('0', prediction=False)
|
mlflow_repository.model_serving.get_model_uri.assert_called_once_with('0', prediction=False)
|
||||||
|
|
||||||
sklearn.load_model.assert_has_calls(
|
sklearn.load_model.assert_has_calls(
|
||||||
|
|||||||
@@ -69,7 +69,7 @@ async def test_run_none_path_flag(workflow_mock, format_and_export_prediction):
|
|||||||
{
|
{
|
||||||
'schema': input_data['schema'],
|
'schema': input_data['schema'],
|
||||||
'table_name': input_data['table_name'],
|
'table_name': input_data['table_name'],
|
||||||
'data': workflow_mock.execute_activity_method.return_value,
|
'data': workflow_mock.execute_local_activity_method.return_value,
|
||||||
**metadata,
|
**metadata,
|
||||||
'timestamp_conversion': {
|
'timestamp_conversion': {
|
||||||
'column': 'timestamp',
|
'column': 'timestamp',
|
||||||
@@ -130,7 +130,7 @@ async def test_run_default_path_flag(workflow_mock, format_and_export_prediction
|
|||||||
{
|
{
|
||||||
'schema': input_data['schema'],
|
'schema': input_data['schema'],
|
||||||
'table_name': input_data['table_name'],
|
'table_name': input_data['table_name'],
|
||||||
'data': workflow_mock.execute_activity_method.return_value,
|
'data': workflow_mock.execute_local_activity_method.return_value,
|
||||||
**metadata,
|
**metadata,
|
||||||
'timestamp_conversion': {
|
'timestamp_conversion': {
|
||||||
'column': 'timestamp',
|
'column': 'timestamp',
|
||||||
@@ -143,5 +143,5 @@ async def test_run_default_path_flag(workflow_mock, format_and_export_prediction
|
|||||||
]
|
]
|
||||||
)
|
)
|
||||||
|
|
||||||
assert workflow_mock.execute_activity_method.call_count == 3
|
assert workflow_mock.execute_activity_method.call_count == 2
|
||||||
assert workflow_mock.execute_local_activity_method.call_count == 1
|
assert workflow_mock.execute_local_activity_method.call_count == 1
|
||||||
|
|||||||
Reference in New Issue
Block a user