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:
vitor-aignosi
2026-06-11 13:34:25 -03:00
parent 19c8a028d2
commit 272e02dadc
3 changed files with 40 additions and 26 deletions

View File

@@ -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:

View File

@@ -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,

View File

@@ -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')