SIENTIAPDE-1255: Implement MLFlow artifact management and model persistence

This commit introduces a new model_repository.py to handle MLFlow artifact generation and model persistence. It also updates the README to reflect this change and modifies training_repository.py to separate training and MLFlow operations.
This commit is contained in:
Bruno Domingues
2025-10-17 09:11:54 -03:00
parent a66b996cc1
commit 9e31ab679b
3 changed files with 18 additions and 916 deletions

View File

@@ -1,9 +1,8 @@
from datetime import UTC, datetime
from unittest.mock import ANY, MagicMock, call, patch
from unittest.mock import MagicMock, patch
import numpy as np
import pytest
from pandas import DataFrame, Timestamp
from pandas import DataFrame
from model_manager.utils.repository.model_repository import MLFlowRepository
@@ -32,477 +31,10 @@ metadata = {
}
class Any:
pass
invalid_cases = [
({'value': {'2024-01-01 12:00:00': 1, 2024: 2}}),
({'value': {'2024-01-01': 1, '2024-01-02': 2}}),
({'value': {Any(): 1, Any(): 2}}),
]
@pytest.mark.parametrize('data', invalid_cases)
def test_detect_and_parse_datetime_index_error_cases(mlflow_repository, data):
input_data = DataFrame(data)
with pytest.raises(ValueError) as e:
mlflow_repository.detect_and_parse_datetime_index(input_data, metadata['metadata'])
assert (
str(e)
== 'Index must be all timestamp like column. Valid formats are: pandas Timestamp, datetime, string in format %Y-%m-%d %H:%M:%S'
)
valid_cases = [
(
{'value': {'2024-01-01 12:00:00+0000': 1, '2024-01-02 12:00:00+0000': 2}},
['2024-01-01 12:00:00+0000', '2024-01-02 12:00:00+0000'],
),
(
{
'value': {
datetime(2025, 1, 1, 12, 0, 0, tzinfo=UTC): 1,
datetime(2025, 1, 2, 12, 0, 0, tzinfo=UTC): 2,
}
},
['2025-01-01 12:00:00+0000', '2025-01-02 12:00:00+0000'],
),
(
{
'value': {
Timestamp(2026, 1, 1, 12, 0, 0, tzinfo=UTC): 1,
Timestamp(2026, 1, 2, 12, 0, 0, tzinfo=UTC): 2,
}
},
['2026-01-01 12:00:00+0000', '2026-01-02 12:00:00+0000'],
),
]
@pytest.mark.parametrize('data,expected', valid_cases)
def test_detect_and_parse_datetime_index_valid_format(mlflow_repository, data, expected):
input_data = DataFrame(data)
response = mlflow_repository.detect_and_parse_datetime_index(input_data, metadata['metadata'])
assert response.index.tolist() == expected
def test_transform_success(mlflow_repository):
data = MagicMock()
model_name = 'model'
mlflow_repository.detect_and_parse_datetime_index = MagicMock()
output = mlflow_repository.transform(model_name, data, {}, metadata['metadata'])
mlflow_repository.model_serving.get_cached_transform.assert_called_once_with(
model_name, data, 0, 'sklearn', False, 'model', 'predict'
)
mlflow_repository.detect_and_parse_datetime_index.assert_called_once_with(
mlflow_repository.model_serving.get_cached_transform.return_value, metadata['metadata']
)
assert output == {
'success': True,
'content': mlflow_repository.detect_and_parse_datetime_index.return_value.to_dict.return_value,
}
def test_transform_error(mlflow_repository):
data = MagicMock()
model_name = 'model'
mlflow_repository.model_serving.get_cached_transform.side_effect = Exception('error')
output = mlflow_repository.transform(model_name, data, {}, metadata['metadata'])
mlflow_repository.model_serving.get_cached_transform.assert_called_once_with(
model_name, data, 0, 'sklearn', False, 'model', 'predict'
)
assert output == {'success': False, 'content': {'message': 'error', 'traceback': ANY}}
def test_predict_success(mlflow_repository):
data = DataFrame({'feat_1': {'index_1': 2, 'index_2': 3}})
model_name = 'model'
mlflow_repository.model_serving.get_cached_predict.return_value = np.array([2, 3])
output = mlflow_repository.predict(model_name, data, {}, metadata['metadata'])
mlflow_repository.model_serving.get_cached_predict.assert_called_once_with(
model_name, data, 0, 'pyfunc', False, 'model'
)
assert output['success'] is True
assert output['content'] == {
'prediction': {'index_1': 2, 'index_2': 3},
'response_time': {'index_1': ANY, 'index_2': ANY},
}
def test_predict_error(mlflow_repository):
data = DataFrame({'feat_1': {'index_1': 2, 'index_2': 3}})
model_name = 'model'
mlflow_repository.model_serving.get_cached_predict = MagicMock(side_effect=Exception('error'))
output = mlflow_repository.predict(model_name, data, {}, metadata['metadata'])
mlflow_repository.model_serving.get_cached_predict.assert_called_once_with(
model_name, data, 0, 'pyfunc', False, 'model'
)
assert output == {'success': False, 'content': {'message': 'error', 'traceback': ANY}}
@patch('model_manager.utils.repository.model_repository.mlflow')
def test_get_experiment_by_run_id(mlflow, mlflow_repository):
mlflow.get_run.return_value = MagicMock(
info=MagicMock(
experiment_id='0',
)
)
mlflow.get_experiment.return_value = MagicMock()
mlflow.get_experiment.return_value.name = 'test'
output = mlflow_repository.get_experiment_by_run_id('0')
assert output == 'test'
mlflow.get_run.assert_called_once_with('0')
mlflow.get_experiment.assert_called_once_with('0')
@patch('model_manager.utils.repository.model_repository.mlflow')
def test_get_next_run_name(mlflow, mlflow_repository):
mlflow.search_runs.return_value = [1, 2, 3]
output = mlflow_repository.get_next_run_name('run')
assert output == 'run-4'
mlflow.search_runs.assert_called_once_with(
experiment_names=['run'],
order_by=['start_time desc'],
)
@patch('model_manager.utils.repository.model_repository.mlflow')
def test_get_experiment_success(mlflow, mlflow_repository):
mlflow.get_experiment_by_name.return_value = MagicMock(experiment_id='0')
output = mlflow_repository.get_experiment('test')
assert output == '0'
@patch('model_manager.utils.repository.model_repository.mlflow')
def test_get_experiment_error(mlflow, mlflow_repository):
mlflow.get_experiment_by_name.return_value = None
try:
mlflow_repository.get_experiment('test')
except ValueError as e:
assert str(e) == 'Experiment test not found'
else:
raise AssertionError('Expected exception')
@patch('model_manager.utils.repository.model_repository.mlflow')
def test_get_experiment_last_run(mlflow, mlflow_repository):
mlflow.search_runs.return_value = DataFrame(
{
'params.retrain': ['True', 'False', 'True', 'False'],
'end_time': ['2021-01-01', '2021-01-02', '2021-01-03', '2021-01-04'],
'run_id': ['0', '1', '2', '3'],
}
)
output = mlflow_repository.get_experiment_last_run(0)
mlflow.search_runs.assert_called_once_with(
experiment_ids=[0],
filter_string='',
output_format='pandas',
)
assert output == '2'
@patch('model_manager.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:
raise AssertionError('Expected exception')
@patch('model_manager.utils.repository.model_repository.mlflow.sklearn')
@patch('model_manager.utils.repository.model_repository.mlflow.set_experiment')
def test_create_model_experiment(set_experiment, sklearn, mlflow_repository):
mlflow_repository.model_serving.get_model_info = MagicMock(return_value='0')
mlflow_repository.model_serving.get_model_uri = MagicMock(return_value='test')
mlflow_repository.get_experiment_by_run_id = MagicMock()
data_model_mock = MagicMock()
prediction_model_mock = MagicMock()
sklearn.load_model.side_effect = [data_model_mock, prediction_model_mock]
data_model_mock.fit.return_value = data_model_mock
data_model_mock.predict.return_value = DataFrame(
{
'x': [10, 20, 30],
}
)
data_model_mock.target_variable = 'y'
prediction_model_mock.fit.return_value = prediction_model_mock
data = DataFrame({'x': [1, 2, 3], 'y': [4, 5, 6]})
output = mlflow_repository.create_model_experiment('test', data)
mlflow_repository.model_serving.get_model_info.assert_called_once_with('test')
mlflow_repository.model_serving.get_model_uri.assert_called_once_with('0', prediction=False)
sklearn.load_model.assert_has_calls(
[
call(mlflow_repository.model_serving.get_model_uri.return_value),
call('models:/test/production'),
]
)
assert sklearn.load_model.call_count == 2
data_model_mock.fit.assert_called_once_with(data)
data_model_mock.predict.assert_called_once_with(data)
fit_args = prediction_model_mock.fit.call_args[0][0]
assert fit_args.equals(
DataFrame(
{
'x': [10, 20, 30],
'y': [4, 5, 6],
}
)
)
mlflow_repository.get_experiment_by_run_id.assert_called_once_with('0')
set_experiment.assert_called_once_with(mlflow_repository.get_experiment_by_run_id.return_value)
assert output == (
prediction_model_mock,
data_model_mock,
mlflow_repository.get_experiment_by_run_id.return_value,
)
@patch('model_manager.utils.repository.model_repository.path.exists')
@patch('model_manager.utils.repository.model_repository.remove')
@patch('model_manager.utils.repository.model_repository.mlflow.start_run')
@patch('model_manager.utils.repository.model_repository.mlflow.log_param')
@patch('model_manager.utils.repository.model_repository.mlflow.sklearn.log_model')
@patch('model_manager.utils.repository.model_repository.mlflow.log_artifact')
def test_perform_model_retrain(
log_artifact, log_model, log_param, start_run, mock_remove, mock_path_exists, mlflow_repository
):
# Create mock models with attributes to test the for loops (lines 268-274)
prediction_model_mock = MagicMock()
prediction_model_mock.__dict__ = {'model': 'pred_model', 'param1': 'value1', 'param2': 'value2'}
data_model_mock = MagicMock()
data_model_mock.__dict__ = {'model': 'data_model', 'param3': 'value3', 'param4': 'value4'}
experiment = 'test'
model_name = 'test'
data = MagicMock()
mlflow_repository.get_next_run_name = MagicMock(return_value='test-1')
run = MagicMock()
start_run.__enter__.return_value = run
mock_path_exists.return_value = True
output = mlflow_repository.perform_model_retrain(
prediction_model_mock, data_model_mock, experiment, model_name, data
)
mlflow_repository.get_next_run_name.assert_called_once_with(experiment)
start_run.assert_called_once_with(
run_name='test-1', description='Retrain model test with new data'
)
log_model.assert_has_calls(
[
call(data_model_mock, 'data_model'),
call(prediction_model_mock, 'prediction_model'),
]
)
data.to_csv.assert_called_once_with('temp/raw_data_test.csv', index=True)
log_artifact.assert_called_once_with('temp/raw_data_test.csv')
# Verify that model attributes were logged (excluding 'model' key)
log_param.assert_has_calls(
[
call('param1', 'value1'), # from prediction_model
call('param2', 'value2'), # from prediction_model
call('param3', 'value3'), # from data_model
call('param4', 'value4'), # from data_model
call('retrain', True),
],
any_order=True,
)
# Verify temp file cleanup
mock_path_exists.assert_called_once_with('temp/raw_data_test.csv')
mock_remove.assert_called_once_with('temp/raw_data_test.csv')
assert output == ('Model retrained successfully', experiment)
@patch('model_manager.utils.repository.model_repository.path.exists')
@patch('model_manager.utils.repository.model_repository.remove')
@patch('model_manager.utils.repository.model_repository.mlflow.start_run')
@patch('model_manager.utils.repository.model_repository.mlflow.log_param')
@patch('model_manager.utils.repository.model_repository.mlflow.sklearn.log_model')
@patch('model_manager.utils.repository.model_repository.mlflow.log_artifact')
def test_perform_model_retrain_file_not_exists(
log_artifact, log_model, log_param, start_run, mock_remove, mock_path_exists, mlflow_repository
):
"""Test perform_model_retrain when temp file doesn't exist (line 291->294 branch)."""
prediction_model_mock = MagicMock()
prediction_model_mock.__dict__ = {'model': 'pred_model'}
data_model_mock = MagicMock()
data_model_mock.__dict__ = {'model': 'data_model'}
experiment = 'test'
model_name = 'test'
data = MagicMock()
mlflow_repository.get_next_run_name = MagicMock(return_value='test-1')
run = MagicMock()
start_run.__enter__.return_value = run
mock_path_exists.return_value = False # File doesn't exist
output = mlflow_repository.perform_model_retrain(
prediction_model_mock, data_model_mock, experiment, model_name, data
)
# Verify temp file cleanup was checked but not executed
mock_path_exists.assert_called_once_with('temp/raw_data_test.csv')
mock_remove.assert_not_called() # Should not be called when file doesn't exist
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('model_manager.utils.repository.model_repository.mlflow')
def test_update_production_model_by_run_id(mlflow, mlflow_repository):
client_mock = MagicMock()
mlflow.tracking.MlflowClient.return_value = client_mock
client_mock.get_registered_model.return_value = MagicMock(
latest_versions=[
MagicMock(version='1'),
MagicMock(version='2'),
MagicMock(version='3'),
]
)
output = mlflow_repository.update_production_model_by_run_id('0', 'test')
mlflow.register_model.assert_called_once_with(
'runs:/0/prediction_model',
'test',
)
mlflow.tracking.MlflowClient.assert_called_once()
client_mock.get_registered_model.assert_called_once_with('test')
client_mock.transition_model_version_stage.assert_called_once_with(
name='test',
version='3',
stage='Production',
archive_existing_versions=True,
)
assert output == {
'model_name': 'test',
'version': '3',
'mlflow_run_id': '0',
}
@patch('model_manager.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: # noqa: BLE001
assert str(e) == 'Model versions is not a list'
else:
raise AssertionError('Expected exception')
def test_update_production_model(mlflow_repository):
connector = mlflow_repository
with patch.object(connector, 'get_experiment', return_value='0') as get_experiment:
with patch.object(
connector, 'get_experiment_last_run', return_value='2'
) as get_experiment_last_run:
with patch.object(
connector,
'update_production_model_by_run_id',
return_value={'model_name': 'test', 'version': '3', 'mlflow_run_id': '0'},
) as update_production_model_by_run_id:
output = connector.update_production_model('0', 'test')
get_experiment.assert_called_once_with('0')
get_experiment_last_run.assert_called_once_with('0')
update_production_model_by_run_id.assert_called_once_with('2', 'test')
assert output == {
'model_name': 'test',
'version': '3',
'mlflow_run_id': '0',
'mlflow_experiment_id': '0',
}
# ========== Tests for Model Artifact Generation Methods ==========
def test_get_next_run_name_new(mlflow_repository):
def test_get_next_run_name(mlflow_repository):
"""Test get_next_run_name generates correct run name based on existing runs."""
mlflow_repository.model_serving.search_runs_by_name.return_value = [
MagicMock(),
@@ -510,7 +42,7 @@ def test_get_next_run_name_new(mlflow_repository):
MagicMock(),
]
result = mlflow_repository.get_next_run_name_new('test_experiment')
result = mlflow_repository.get_next_run_name('test_experiment')
mlflow_repository.model_serving.search_runs_by_name.assert_called_once_with(
experiment_names=['test_experiment'], order_by=['start_time desc']
@@ -518,13 +50,13 @@ def test_get_next_run_name_new(mlflow_repository):
assert result == 'test_experiment-4'
def test_get_next_run_name_new_first_run(mlflow_repository):
def test_get_next_run_name_first_run(mlflow_repository):
"""Test get_next_run_name for first run (no existing runs)."""
mlflow_repository.model_serving.search_runs_by_name.return_value = []
result = mlflow_repository.get_next_run_name_new('new_experiment')
result = mlflow_repository.get_next_run_name('test_experiment')
assert result == 'new_experiment-1'
assert result == 'test_experiment-1'
@patch('model_manager.utils.repository.model_repository.path')