SIENTIAPDE-1646
SIENTIAPDE-1646 Refactor MLFlow tests and update artifact handling - Enhanced test cases for MLFlow to improve clarity and accuracy in data handling. - Updated references in tests to use 'evaluation_data.csv' and 'test_data.csv' instead of 'retrain_input.csv' and 'train_data.csv'. - Introduced a new helper function for creating prediction frames to streamline test setup.
This commit is contained in:
@@ -559,7 +559,7 @@ class MLFlow(SientiaMonitoring):
|
|||||||
evaluation_csv = Path(tmp_dir) / 'evaluation_data.csv'
|
evaluation_csv = Path(tmp_dir) / 'evaluation_data.csv'
|
||||||
data.to_csv(raw_csv, index=False)
|
data.to_csv(raw_csv, index=False)
|
||||||
evaluation_data.to_csv(evaluation_csv, index=False)
|
evaluation_data.to_csv(evaluation_csv, index=False)
|
||||||
|
|
||||||
mlflow.log_artifact(str(raw_csv))
|
mlflow.log_artifact(str(raw_csv))
|
||||||
mlflow.log_artifact(str(evaluation_csv))
|
mlflow.log_artifact(str(evaluation_csv))
|
||||||
finally:
|
finally:
|
||||||
|
|||||||
@@ -216,7 +216,7 @@ async def main():
|
|||||||
activities.export_data_to_postgres,
|
activities.export_data_to_postgres,
|
||||||
],
|
],
|
||||||
logger=logger,
|
logger=logger,
|
||||||
runtime='core'
|
runtime='core',
|
||||||
),
|
),
|
||||||
prepare_worker(
|
prepare_worker(
|
||||||
temporal_client=temporal_client,
|
temporal_client=temporal_client,
|
||||||
@@ -229,7 +229,7 @@ async def main():
|
|||||||
activities.export_data_to_postgres,
|
activities.export_data_to_postgres,
|
||||||
],
|
],
|
||||||
logger=logger,
|
logger=logger,
|
||||||
runtime='core'
|
runtime='core',
|
||||||
),
|
),
|
||||||
prepare_worker(
|
prepare_worker(
|
||||||
temporal_client=temporal_client,
|
temporal_client=temporal_client,
|
||||||
|
|||||||
@@ -399,6 +399,8 @@ def test_retrain_model_success_data_success_retrain(
|
|||||||
'value': [1.0, 2.0],
|
'value': [1.0, 2.0],
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
pivoted_index = pd.DatetimeIndex([ts])
|
||||||
|
wrapper.retrain.return_value = _retrain_prediction_frame(pivoted_index, 'target', [1.0])
|
||||||
|
|
||||||
payload = MagicMock()
|
payload = MagicMock()
|
||||||
payload.retrieve = MagicMock(return_value=raw_data)
|
payload.retrieve = MagicMock(return_value=raw_data)
|
||||||
@@ -428,6 +430,7 @@ def test_retrain_model_success_with_payload_data(
|
|||||||
mock_to_datetime, mock_rmtree, mock_mkdtemp, mock_log_artifact, mlflow
|
mock_to_datetime, mock_rmtree, mock_mkdtemp, mock_log_artifact, mlflow
|
||||||
):
|
):
|
||||||
mock_mkdtemp.return_value = 'tmp'
|
mock_mkdtemp.return_value = 'tmp'
|
||||||
|
mock_to_datetime.side_effect = lambda idx, **kwargs: pd.DatetimeIndex(idx)
|
||||||
mv_alias = MagicMock(run_id='src')
|
mv_alias = MagicMock(run_id='src')
|
||||||
mlflow.mlflow_repository._client.get_model_version_by_alias.return_value = mv_alias
|
mlflow.mlflow_repository._client.get_model_version_by_alias.return_value = mv_alias
|
||||||
wrapper = MagicMock()
|
wrapper = MagicMock()
|
||||||
@@ -437,20 +440,21 @@ def test_retrain_model_success_with_payload_data(
|
|||||||
mock_cm.__exit__.return_value = False
|
mock_cm.__exit__.return_value = False
|
||||||
mlflow.mlflow_repository.start_run.return_value = mock_cm
|
mlflow.mlflow_repository.start_run.return_value = mock_cm
|
||||||
|
|
||||||
raw_data = MagicMock(columns=['variable', 'timestamp', 'value', 'created_at'])
|
ts = pd.Timestamp('2020-01-01', tz='UTC')
|
||||||
raw_data.__getitem__.return_value.max.return_value = 'ts'
|
raw_data = pd.DataFrame(
|
||||||
|
{
|
||||||
|
'variable': ['target', 'f1', 'target', 'f1'],
|
||||||
|
'timestamp': [ts, ts, ts + pd.Timedelta(hours=1), ts + pd.Timedelta(hours=1)],
|
||||||
|
'value': [1.0, 2.0, 3.0, 4.0],
|
||||||
|
'created_at': [ts, ts, ts + pd.Timedelta(hours=1), ts + pd.Timedelta(hours=1)],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
pivoted_index = pd.DatetimeIndex([ts, ts + pd.Timedelta(hours=1)])
|
||||||
|
wrapper.retrain.return_value = _retrain_prediction_frame(pivoted_index, 'target', [1.0, 3.0])
|
||||||
|
|
||||||
payload = MagicMock()
|
payload = MagicMock()
|
||||||
payload.retrieve = MagicMock(return_value=raw_data)
|
payload.retrieve = MagicMock(return_value=raw_data)
|
||||||
|
|
||||||
pivoted = MagicMock()
|
|
||||||
raw_data.sort_values.return_value = raw_data
|
|
||||||
raw_data.drop_duplicates.return_value = raw_data
|
|
||||||
raw_data.pivot.return_value = pivoted
|
|
||||||
pivoted.fillna = MagicMock()
|
|
||||||
pivoted.columns.name = None
|
|
||||||
pivoted.index = MagicMock()
|
|
||||||
pivoted.__setitem__ = MagicMock()
|
|
||||||
|
|
||||||
response = mlflow.retrain_model(
|
response = mlflow.retrain_model(
|
||||||
{
|
{
|
||||||
**metadata,
|
**metadata,
|
||||||
@@ -482,13 +486,17 @@ def test_retrain_model_always_uses_retrain_even_with_full_retrain_flag(
|
|||||||
mlflow.mlflow_repository.start_run.return_value = mock_cm
|
mlflow.mlflow_repository.start_run.return_value = mock_cm
|
||||||
|
|
||||||
ts = pd.Timestamp('2020-01-01', tz='UTC')
|
ts = pd.Timestamp('2020-01-01', tz='UTC')
|
||||||
|
ts_next = ts + pd.Timedelta(days=1)
|
||||||
raw_data = pd.DataFrame(
|
raw_data = pd.DataFrame(
|
||||||
{
|
{
|
||||||
'variable': ['target', 'f1', 'target', 'f1'],
|
'variable': ['target', 'f1', 'target', 'f1'],
|
||||||
'timestamp': [ts, ts, ts + pd.Timedelta(days=1), ts + pd.Timedelta(days=1)],
|
'timestamp': [ts, ts, ts_next, ts_next],
|
||||||
'value': [1.0, 2.0, 3.0, 4.0],
|
'value': [1.0, 2.0, 3.0, 4.0],
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
pivoted_index = pd.DatetimeIndex([ts, ts_next])
|
||||||
|
wrapper.retrain.return_value = _retrain_prediction_frame(pivoted_index, 'target', [1.0, 3.0])
|
||||||
|
|
||||||
payload = MagicMock()
|
payload = MagicMock()
|
||||||
payload.retrieve = MagicMock(return_value=raw_data)
|
payload.retrieve = MagicMock(return_value=raw_data)
|
||||||
|
|
||||||
@@ -658,6 +666,12 @@ def _artifact_file_info(path: str) -> MagicMock:
|
|||||||
return file_info
|
return file_info
|
||||||
|
|
||||||
|
|
||||||
|
def _retrain_prediction_frame(
|
||||||
|
index: pd.DatetimeIndex, target: str, values: list[float]
|
||||||
|
) -> pd.DataFrame:
|
||||||
|
return pd.DataFrame({target: values}, index=index)
|
||||||
|
|
||||||
|
|
||||||
@patch('laborious.activities.mlflow.to_datetime')
|
@patch('laborious.activities.mlflow.to_datetime')
|
||||||
def test_get_reference_data_success(mock_to_datetime, mlflow):
|
def test_get_reference_data_success(mock_to_datetime, mlflow):
|
||||||
input_data = {
|
input_data = {
|
||||||
@@ -668,7 +682,7 @@ def test_get_reference_data_success(mock_to_datetime, mlflow):
|
|||||||
mv = MagicMock(run_id='run1')
|
mv = MagicMock(run_id='run1')
|
||||||
mlflow.mlflow_repository._client.get_model_version_by_alias.return_value = mv
|
mlflow.mlflow_repository._client.get_model_version_by_alias.return_value = mv
|
||||||
mlflow.mlflow_repository._client.list_artifacts.return_value = [
|
mlflow.mlflow_repository._client.list_artifacts.return_value = [
|
||||||
_artifact_file_info('retrain_input.csv'),
|
_artifact_file_info('evaluation_data.csv'),
|
||||||
]
|
]
|
||||||
|
|
||||||
mock_reference_data = MagicMock()
|
mock_reference_data = MagicMock()
|
||||||
@@ -690,7 +704,7 @@ def test_get_reference_data_success(mock_to_datetime, mlflow):
|
|||||||
|
|
||||||
mlflow.mlflow_repository.download_artifacts.assert_called_once_with(
|
mlflow.mlflow_repository.download_artifacts.assert_called_once_with(
|
||||||
run_id='run1',
|
run_id='run1',
|
||||||
artifact_path='retrain_input.csv',
|
artifact_path='evaluation_data.csv',
|
||||||
dst_path='/t',
|
dst_path='/t',
|
||||||
metadata=metadata['metadata'],
|
metadata=metadata['metadata'],
|
||||||
)
|
)
|
||||||
@@ -712,7 +726,7 @@ def test_get_reference_data_not_found(mlflow):
|
|||||||
assert result is None
|
assert result is None
|
||||||
|
|
||||||
|
|
||||||
def test_get_reference_data_only_train_data_csv(mlflow):
|
def test_get_reference_data_only_test_data_csv(mlflow):
|
||||||
input_data = {
|
input_data = {
|
||||||
**metadata,
|
**metadata,
|
||||||
'model_name': 'test_model',
|
'model_name': 'test_model',
|
||||||
@@ -720,7 +734,7 @@ def test_get_reference_data_only_train_data_csv(mlflow):
|
|||||||
mv = MagicMock(run_id='run1')
|
mv = MagicMock(run_id='run1')
|
||||||
mlflow.mlflow_repository._client.get_model_version_by_alias.return_value = mv
|
mlflow.mlflow_repository._client.get_model_version_by_alias.return_value = mv
|
||||||
mlflow.mlflow_repository._client.list_artifacts.return_value = [
|
mlflow.mlflow_repository._client.list_artifacts.return_value = [
|
||||||
_artifact_file_info('train_data.csv'),
|
_artifact_file_info('test_data.csv'),
|
||||||
]
|
]
|
||||||
|
|
||||||
mock_reference_data = MagicMock()
|
mock_reference_data = MagicMock()
|
||||||
@@ -741,14 +755,14 @@ def test_get_reference_data_only_train_data_csv(mlflow):
|
|||||||
|
|
||||||
mlflow.mlflow_repository.download_artifacts.assert_called_once_with(
|
mlflow.mlflow_repository.download_artifacts.assert_called_once_with(
|
||||||
run_id='run1',
|
run_id='run1',
|
||||||
artifact_path='train_data.csv',
|
artifact_path='test_data.csv',
|
||||||
dst_path='/t',
|
dst_path='/t',
|
||||||
metadata=metadata['metadata'],
|
metadata=metadata['metadata'],
|
||||||
)
|
)
|
||||||
assert result == mock_reference_data.to_dict.return_value
|
assert result == mock_reference_data.to_dict.return_value
|
||||||
|
|
||||||
|
|
||||||
def test_get_reference_data_prefers_retrain_input_when_both_listed(mlflow):
|
def test_get_reference_data_prefers_evaluation_data_when_both_listed(mlflow):
|
||||||
input_data = {
|
input_data = {
|
||||||
**metadata,
|
**metadata,
|
||||||
'model_name': 'test_model',
|
'model_name': 'test_model',
|
||||||
@@ -756,8 +770,8 @@ def test_get_reference_data_prefers_retrain_input_when_both_listed(mlflow):
|
|||||||
mv = MagicMock(run_id='run1')
|
mv = MagicMock(run_id='run1')
|
||||||
mlflow.mlflow_repository._client.get_model_version_by_alias.return_value = mv
|
mlflow.mlflow_repository._client.get_model_version_by_alias.return_value = mv
|
||||||
mlflow.mlflow_repository._client.list_artifacts.return_value = [
|
mlflow.mlflow_repository._client.list_artifacts.return_value = [
|
||||||
_artifact_file_info('train_data.csv'),
|
_artifact_file_info('test_data.csv'),
|
||||||
_artifact_file_info('retrain_input.csv'),
|
_artifact_file_info('evaluation_data.csv'),
|
||||||
]
|
]
|
||||||
|
|
||||||
mock_reference_data = MagicMock()
|
mock_reference_data = MagicMock()
|
||||||
@@ -778,7 +792,7 @@ def test_get_reference_data_prefers_retrain_input_when_both_listed(mlflow):
|
|||||||
|
|
||||||
mlflow.mlflow_repository.download_artifacts.assert_called_once_with(
|
mlflow.mlflow_repository.download_artifacts.assert_called_once_with(
|
||||||
run_id='run1',
|
run_id='run1',
|
||||||
artifact_path='retrain_input.csv',
|
artifact_path='evaluation_data.csv',
|
||||||
dst_path='/t',
|
dst_path='/t',
|
||||||
metadata=metadata['metadata'],
|
metadata=metadata['metadata'],
|
||||||
)
|
)
|
||||||
@@ -811,7 +825,7 @@ def test_get_reference_data_missing_csv_file_returns_none(mlflow):
|
|||||||
mv = MagicMock(run_id='run1')
|
mv = MagicMock(run_id='run1')
|
||||||
mlflow.mlflow_repository._client.get_model_version_by_alias.return_value = mv
|
mlflow.mlflow_repository._client.get_model_version_by_alias.return_value = mv
|
||||||
mlflow.mlflow_repository._client.list_artifacts.return_value = [
|
mlflow.mlflow_repository._client.list_artifacts.return_value = [
|
||||||
_artifact_file_info('retrain_input.csv'),
|
_artifact_file_info('evaluation_data.csv'),
|
||||||
]
|
]
|
||||||
|
|
||||||
with patch('laborious.activities.mlflow.tempfile.mkdtemp', return_value='tmp'):
|
with patch('laborious.activities.mlflow.tempfile.mkdtemp', return_value='tmp'):
|
||||||
@@ -833,7 +847,7 @@ def test_get_reference_data_exception(mlflow):
|
|||||||
mv = MagicMock(run_id='run1')
|
mv = MagicMock(run_id='run1')
|
||||||
mlflow.mlflow_repository._client.get_model_version_by_alias.return_value = mv
|
mlflow.mlflow_repository._client.get_model_version_by_alias.return_value = mv
|
||||||
mlflow.mlflow_repository._client.list_artifacts.return_value = [
|
mlflow.mlflow_repository._client.list_artifacts.return_value = [
|
||||||
_artifact_file_info('retrain_input.csv'),
|
_artifact_file_info('evaluation_data.csv'),
|
||||||
]
|
]
|
||||||
mlflow.mlflow_repository.download_artifacts.side_effect = Exception('dl fail')
|
mlflow.mlflow_repository.download_artifacts.side_effect = Exception('dl fail')
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user