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:
@@ -216,7 +216,7 @@ async def main():
|
||||
activities.export_data_to_postgres,
|
||||
],
|
||||
logger=logger,
|
||||
runtime='core'
|
||||
runtime='core',
|
||||
),
|
||||
prepare_worker(
|
||||
temporal_client=temporal_client,
|
||||
@@ -229,7 +229,7 @@ async def main():
|
||||
activities.export_data_to_postgres,
|
||||
],
|
||||
logger=logger,
|
||||
runtime='core'
|
||||
runtime='core',
|
||||
),
|
||||
prepare_worker(
|
||||
temporal_client=temporal_client,
|
||||
|
||||
@@ -399,6 +399,8 @@ def test_retrain_model_success_data_success_retrain(
|
||||
'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.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_mkdtemp.return_value = 'tmp'
|
||||
mock_to_datetime.side_effect = lambda idx, **kwargs: pd.DatetimeIndex(idx)
|
||||
mv_alias = MagicMock(run_id='src')
|
||||
mlflow.mlflow_repository._client.get_model_version_by_alias.return_value = mv_alias
|
||||
wrapper = MagicMock()
|
||||
@@ -437,20 +440,21 @@ def test_retrain_model_success_with_payload_data(
|
||||
mock_cm.__exit__.return_value = False
|
||||
mlflow.mlflow_repository.start_run.return_value = mock_cm
|
||||
|
||||
raw_data = MagicMock(columns=['variable', 'timestamp', 'value', 'created_at'])
|
||||
raw_data.__getitem__.return_value.max.return_value = 'ts'
|
||||
ts = pd.Timestamp('2020-01-01', tz='UTC')
|
||||
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.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(
|
||||
{
|
||||
**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
|
||||
|
||||
ts = pd.Timestamp('2020-01-01', tz='UTC')
|
||||
ts_next = ts + pd.Timedelta(days=1)
|
||||
raw_data = pd.DataFrame(
|
||||
{
|
||||
'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],
|
||||
}
|
||||
)
|
||||
pivoted_index = pd.DatetimeIndex([ts, ts_next])
|
||||
wrapper.retrain.return_value = _retrain_prediction_frame(pivoted_index, 'target', [1.0, 3.0])
|
||||
|
||||
payload = MagicMock()
|
||||
payload.retrieve = MagicMock(return_value=raw_data)
|
||||
|
||||
@@ -658,6 +666,12 @@ def _artifact_file_info(path: str) -> MagicMock:
|
||||
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')
|
||||
def test_get_reference_data_success(mock_to_datetime, mlflow):
|
||||
input_data = {
|
||||
@@ -668,7 +682,7 @@ def test_get_reference_data_success(mock_to_datetime, mlflow):
|
||||
mv = MagicMock(run_id='run1')
|
||||
mlflow.mlflow_repository._client.get_model_version_by_alias.return_value = mv
|
||||
mlflow.mlflow_repository._client.list_artifacts.return_value = [
|
||||
_artifact_file_info('retrain_input.csv'),
|
||||
_artifact_file_info('evaluation_data.csv'),
|
||||
]
|
||||
|
||||
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(
|
||||
run_id='run1',
|
||||
artifact_path='retrain_input.csv',
|
||||
artifact_path='evaluation_data.csv',
|
||||
dst_path='/t',
|
||||
metadata=metadata['metadata'],
|
||||
)
|
||||
@@ -712,7 +726,7 @@ def test_get_reference_data_not_found(mlflow):
|
||||
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 = {
|
||||
**metadata,
|
||||
'model_name': 'test_model',
|
||||
@@ -720,7 +734,7 @@ def test_get_reference_data_only_train_data_csv(mlflow):
|
||||
mv = MagicMock(run_id='run1')
|
||||
mlflow.mlflow_repository._client.get_model_version_by_alias.return_value = mv
|
||||
mlflow.mlflow_repository._client.list_artifacts.return_value = [
|
||||
_artifact_file_info('train_data.csv'),
|
||||
_artifact_file_info('test_data.csv'),
|
||||
]
|
||||
|
||||
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(
|
||||
run_id='run1',
|
||||
artifact_path='train_data.csv',
|
||||
artifact_path='test_data.csv',
|
||||
dst_path='/t',
|
||||
metadata=metadata['metadata'],
|
||||
)
|
||||
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 = {
|
||||
**metadata,
|
||||
'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')
|
||||
mlflow.mlflow_repository._client.get_model_version_by_alias.return_value = mv
|
||||
mlflow.mlflow_repository._client.list_artifacts.return_value = [
|
||||
_artifact_file_info('train_data.csv'),
|
||||
_artifact_file_info('retrain_input.csv'),
|
||||
_artifact_file_info('test_data.csv'),
|
||||
_artifact_file_info('evaluation_data.csv'),
|
||||
]
|
||||
|
||||
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(
|
||||
run_id='run1',
|
||||
artifact_path='retrain_input.csv',
|
||||
artifact_path='evaluation_data.csv',
|
||||
dst_path='/t',
|
||||
metadata=metadata['metadata'],
|
||||
)
|
||||
@@ -811,7 +825,7 @@ def test_get_reference_data_missing_csv_file_returns_none(mlflow):
|
||||
mv = MagicMock(run_id='run1')
|
||||
mlflow.mlflow_repository._client.get_model_version_by_alias.return_value = mv
|
||||
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'):
|
||||
@@ -833,7 +847,7 @@ def test_get_reference_data_exception(mlflow):
|
||||
mv = MagicMock(run_id='run1')
|
||||
mlflow.mlflow_repository._client.get_model_version_by_alias.return_value = mv
|
||||
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')
|
||||
|
||||
|
||||
Reference in New Issue
Block a user