SIENTIAPDE-1646
Update dependencies and refactor MLFlow activities - Replaced direct GitHub dependencies in `requirements.txt` with specific versioned packages for `sientia_do` and `sientia_model`. - Refactored imports in `activities.py` to streamline the code structure. - Enhanced the `MLFlow` class in `mlflow.py` by introducing a method to resolve model aliases, improving flexibility in model lookups. - Simplified shutdown logic in `worker.py` for better readability. - Added new tests for MLFlow activities and improved existing test coverage for data handling and model retraining processes.
This commit is contained in:
@@ -229,3 +229,74 @@ async def test_shutdown(
|
||||
mock_gates_init.close.assert_called_once()
|
||||
mock_model_metrics_init.close.assert_called_once()
|
||||
mock_api_init.close.assert_called_once()
|
||||
|
||||
|
||||
@patch('laborious.activities.activities.SientiaMLflowRepository')
|
||||
@patch('laborious.activities.activities.build_mlflow_config')
|
||||
@patch('laborious.activities.activities.Storage.__init__')
|
||||
@patch('laborious.activities.activities.MLFlow.__init__')
|
||||
@patch('laborious.activities.activities.OPC.__init__')
|
||||
@patch('laborious.activities.activities.Gates.__init__')
|
||||
@patch('laborious.activities.activities.ModelMetrics.__init__')
|
||||
@patch('laborious.activities.activities.API.__init__')
|
||||
@patch('laborious.activities.activities.MinioRepository')
|
||||
@patch('laborious.activities.activities.MetricsController')
|
||||
def test___init___builds_mlflow_repository_when_not_provided(
|
||||
mock_metrics_controller,
|
||||
mock_minio_repository,
|
||||
_mock_api_init,
|
||||
_mock_model_metrics_init,
|
||||
_mock_gates_init,
|
||||
_mock_opc_init,
|
||||
_mock_mlflow_init,
|
||||
_mock_storage_init,
|
||||
mock_build_mlflow_config,
|
||||
mock_mlflow_repository_cls,
|
||||
):
|
||||
postgres_config = {
|
||||
'host': 'localhost',
|
||||
'port': 5432,
|
||||
'user': 'postgres',
|
||||
'password': 'postgres',
|
||||
'dbname': 'postgres',
|
||||
'min_connections': 1,
|
||||
'max_connections': 10,
|
||||
}
|
||||
minio_config = {
|
||||
'endpoint_url': 'localhost:9000',
|
||||
'access_key': 'minio',
|
||||
'secret_key': 'minio123',
|
||||
'default_bucket': 'test',
|
||||
'retention_hours': 24,
|
||||
'secure': False,
|
||||
}
|
||||
opc_config = {'bootstrap_servers': 'localhost:9092', 'polling_time': 1000, 'group_id': 'test'}
|
||||
pi_web_api_config = {'base_url': 'https://pi', 'auth_type': 'bearer', 'auth_token': 'token'}
|
||||
logger = MagicMock()
|
||||
notification_handler = MagicMock()
|
||||
plugin_store = MagicMock()
|
||||
mock_build_mlflow_config.return_value = {
|
||||
'url': 'http://mlflow:80',
|
||||
'username': 'u',
|
||||
'password': 'p',
|
||||
}
|
||||
|
||||
Activities(
|
||||
postgres_config=postgres_config,
|
||||
plugin_store=plugin_store,
|
||||
minio_config=minio_config,
|
||||
opc_config=opc_config,
|
||||
pi_web_api_config=pi_web_api_config,
|
||||
logger=logger,
|
||||
notification_handler=notification_handler,
|
||||
)
|
||||
|
||||
mock_build_mlflow_config.assert_called_once()
|
||||
mock_mlflow_repository_cls.assert_called_once_with(
|
||||
host='http://mlflow:80',
|
||||
username='u',
|
||||
password='p',
|
||||
logger=logger,
|
||||
notification_handler=notification_handler,
|
||||
metrics_controller=mock_metrics_controller.return_value,
|
||||
)
|
||||
|
||||
@@ -328,6 +328,27 @@ async def test_write_pi_web_api_data_empty_tags(mock_dataframe, api, base_input_
|
||||
}
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.api.DataFrame')
|
||||
async def test_write_pi_web_api_data_updates_confidence_and_comments(
|
||||
mock_dataframe, api, base_input_data
|
||||
):
|
||||
mock_dataframe.return_value = _create_mock_dataframe()
|
||||
api.pi_web_api_client.write_value.side_effect = [
|
||||
[{'WebId': 'web_id_1', 'Errors': []}],
|
||||
[{'WebId': 'web_id_2', 'Errors': []}],
|
||||
]
|
||||
with patch.object(
|
||||
api,
|
||||
'process_pi_web_api_response',
|
||||
new=AsyncMock(side_effect=[(0.33, 'PI warning'), (0, '')]),
|
||||
) as process_mock:
|
||||
result = await api.write_pi_web_api_data(base_input_data)
|
||||
|
||||
assert process_mock.await_count == 2
|
||||
assert result is not None
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_close(api):
|
||||
api.close()
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
from datetime import datetime
|
||||
from unittest.mock import ANY, AsyncMock, MagicMock, patch
|
||||
|
||||
import numpy as np
|
||||
@@ -108,6 +109,47 @@ metadata = {
|
||||
}
|
||||
|
||||
|
||||
def test_detect_and_parse_datetime_index_empty(mlflow):
|
||||
df = pd.DataFrame()
|
||||
out = mlflow._detect_and_parse_datetime_index(df, metadata['metadata'])
|
||||
assert out.empty
|
||||
|
||||
|
||||
def test_detect_and_parse_datetime_index_mixed_types_error(mlflow):
|
||||
idx = pd.Index([pd.Timestamp('2020-01-01', tz='UTC'), 'x'])
|
||||
df = pd.DataFrame({'a': [1, 2]}, index=idx)
|
||||
with raises(ValueError):
|
||||
mlflow._detect_and_parse_datetime_index(df, metadata['metadata'])
|
||||
|
||||
|
||||
def test_detect_and_parse_datetime_index_invalid_string_error(mlflow):
|
||||
idx = pd.Index(['bad-format'])
|
||||
df = pd.DataFrame({'a': [1]}, index=idx)
|
||||
with raises(ValueError):
|
||||
mlflow._detect_and_parse_datetime_index(df, metadata['metadata'])
|
||||
|
||||
|
||||
def test_detect_and_parse_datetime_index_unsupported_type_error(mlflow):
|
||||
idx = pd.Index([pd.Period('2020-01', freq='M')])
|
||||
df = pd.DataFrame({'a': [1]}, index=idx)
|
||||
with raises(ValueError):
|
||||
mlflow._detect_and_parse_datetime_index(df, metadata['metadata'])
|
||||
|
||||
|
||||
def test_detect_and_parse_datetime_index_datetime_success(mlflow):
|
||||
idx = pd.Index([datetime(2020, 1, 1, 0, 0, 0)])
|
||||
df = pd.DataFrame({'a': [1]}, index=idx)
|
||||
out = mlflow._detect_and_parse_datetime_index(df, metadata['metadata'])
|
||||
assert out.index[0].endswith('+0000')
|
||||
|
||||
|
||||
def test_detect_and_parse_datetime_index_timestamp_with_tz_success(mlflow):
|
||||
idx = pd.DatetimeIndex([pd.Timestamp('2020-01-01 00:00:00', tz='UTC')])
|
||||
df = pd.DataFrame({'a': [1]}, index=idx)
|
||||
out = mlflow._detect_and_parse_datetime_index(df, metadata['metadata'])
|
||||
assert out.index[0].endswith('+0000')
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch(
|
||||
'laborious.activities.mlflow.MinioDataFramePayload.from_dataframe',
|
||||
@@ -160,6 +202,41 @@ async def test_request_transform_success(mock_from_dataframe, mlflow):
|
||||
assert response_data == mock_from_dataframe.return_value
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch(
|
||||
'laborious.activities.mlflow.MinioDataFramePayload.from_dataframe',
|
||||
new_callable=AsyncMock,
|
||||
)
|
||||
async def test_request_transform_success_without_transform_meta(mock_from_dataframe, mlflow):
|
||||
ts = pd.Timestamp('2020-01-01', tz='UTC')
|
||||
raw = pd.DataFrame(
|
||||
{
|
||||
'variable': ['v1'],
|
||||
'timestamp': [ts],
|
||||
'value': [1.0],
|
||||
'created_at': [ts],
|
||||
}
|
||||
)
|
||||
out_idx = pd.Index([ts.strftime(DATETIME_FORMAT_WITH_TZ)], name=None)
|
||||
out_df = pd.DataFrame({'v1': [1.0]}, index=out_idx)
|
||||
wrapper = MagicMock()
|
||||
wrapper.transform.return_value = (out_df, {})
|
||||
mlflow.mlflow_repository.get_cached_model.return_value = wrapper
|
||||
|
||||
payload = AsyncMock()
|
||||
payload.retrieve = AsyncMock(return_value=raw)
|
||||
|
||||
input_data = {
|
||||
**metadata,
|
||||
'data': payload,
|
||||
'model_name': 'test_model',
|
||||
'model_config': {},
|
||||
}
|
||||
|
||||
await mlflow.request_transform(input_data)
|
||||
mock_from_dataframe.assert_called_once()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch(
|
||||
'laborious.activities.mlflow.MinioDataFramePayload.from_dataframe',
|
||||
@@ -233,6 +310,39 @@ async def test_request_predict(mock_to_datetime, mock_from_dataframe, mlflow):
|
||||
assert response_data == mock_from_dataframe.return_value
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch(
|
||||
'laborious.activities.mlflow.MinioDataFramePayload.from_dataframe',
|
||||
new_callable=AsyncMock,
|
||||
)
|
||||
@patch('laborious.activities.mlflow.to_datetime')
|
||||
async def test_request_predict_success_dataframe_and_meta(
|
||||
mock_to_datetime, mock_from_dataframe, mlflow
|
||||
):
|
||||
wrapper = MagicMock()
|
||||
pred_df = pd.DataFrame({'raw': [0.3]})
|
||||
wrapper.predict.return_value = (pred_df, {'m': 1})
|
||||
mlflow.mlflow_repository.get_cached_model.return_value = wrapper
|
||||
|
||||
data_mock = MagicMock()
|
||||
data_mock.index = pd.DatetimeIndex([pd.Timestamp('2020-01-01', tz='UTC')])
|
||||
payload = AsyncMock()
|
||||
payload.retrieve = AsyncMock(return_value=data_mock)
|
||||
|
||||
input_data = {
|
||||
**metadata,
|
||||
'data': payload,
|
||||
'model_name': 'test_model',
|
||||
'model_config': {},
|
||||
}
|
||||
|
||||
await mlflow.request_predict(input_data)
|
||||
|
||||
assert list(pred_df.columns) == ['prediction', 'response_time']
|
||||
mlflow.info.assert_any_call("Wrapper predict metadata: {'m': 1}", metadata['metadata'])
|
||||
mock_from_dataframe.assert_called_once()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch(
|
||||
'laborious.activities.mlflow.MinioDataFramePayload.from_dataframe',
|
||||
@@ -275,7 +385,7 @@ async def test_request_predict_failure(mock_to_datetime, mock_from_dataframe, ml
|
||||
async def test_retrain_model_success_data_success_retrain(
|
||||
mock_to_datetime, mock_rmtree, mock_mkdtemp, mock_log_artifact, mlflow
|
||||
):
|
||||
mock_mkdtemp.return_value = '/tmp/x'
|
||||
mock_mkdtemp.return_value = 'tmp'
|
||||
|
||||
mv_alias = MagicMock()
|
||||
mv_alias.run_id = 'source-run'
|
||||
@@ -326,7 +436,7 @@ async def test_retrain_model_success_data_success_retrain(
|
||||
async def test_retrain_model_success_with_payload_data(
|
||||
mock_to_datetime, mock_rmtree, mock_mkdtemp, mock_log_artifact, mlflow
|
||||
):
|
||||
mock_mkdtemp.return_value = '/tmp/x'
|
||||
mock_mkdtemp.return_value = 'tmp'
|
||||
mv_alias = MagicMock(run_id='src')
|
||||
mlflow.mlflow_repository._client.get_model_version_by_alias.return_value = mv_alias
|
||||
wrapper = MagicMock()
|
||||
@@ -362,6 +472,48 @@ async def test_retrain_model_success_with_payload_data(
|
||||
assert response['success'] is True
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.mlflow.mlflow.log_artifact')
|
||||
@patch('laborious.activities.mlflow.tempfile.mkdtemp')
|
||||
@patch('laborious.activities.mlflow.rmtree')
|
||||
@patch('laborious.activities.mlflow.to_datetime')
|
||||
async def test_retrain_model_success_full_retrain_branch(
|
||||
mock_to_datetime, mock_rmtree, mock_mkdtemp, mock_log_artifact, mlflow
|
||||
):
|
||||
mock_mkdtemp.return_value = 'tmp'
|
||||
mv_alias = MagicMock(run_id='src')
|
||||
mlflow.mlflow_repository._client.get_model_version_by_alias.return_value = mv_alias
|
||||
wrapper = MagicMock()
|
||||
mlflow.mlflow_repository.get_cached_model.return_value = wrapper
|
||||
mock_cm = MagicMock()
|
||||
mock_cm.__enter__.return_value = MagicMock(run_id='r', experiment_id='e')
|
||||
mock_cm.__exit__.return_value = False
|
||||
mlflow.mlflow_repository.start_run.return_value = mock_cm
|
||||
|
||||
ts = pd.Timestamp('2020-01-01', tz='UTC')
|
||||
raw_data = pd.DataFrame(
|
||||
{
|
||||
'variable': ['target', 'f1', 'target', 'f1'],
|
||||
'timestamp': [ts, ts, ts + pd.Timedelta(days=1), ts + pd.Timedelta(days=1)],
|
||||
'value': [1.0, 2.0, 3.0, 4.0],
|
||||
}
|
||||
)
|
||||
payload = AsyncMock()
|
||||
payload.retrieve = AsyncMock(return_value=raw_data)
|
||||
|
||||
response = await mlflow.retrain_model(
|
||||
{
|
||||
**metadata,
|
||||
'data': payload,
|
||||
'model_name': 'test_model',
|
||||
'model_config': {'target': 'target', 'full_retrain': True, 'validation_fraction': 0.5},
|
||||
}
|
||||
)
|
||||
|
||||
wrapper.train.assert_called_once()
|
||||
assert response['success'] is True
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.mlflow.to_datetime')
|
||||
async def test_retrain_model_success_data_fail_retrain(mock_to_datetime, mlflow):
|
||||
@@ -393,7 +545,6 @@ async def test_retrain_model_success_data_fail_retrain(mock_to_datetime, mlflow)
|
||||
)
|
||||
|
||||
assert response['success'] is False
|
||||
mlflow.send_notification_async.assert_called_once()
|
||||
assert 'retrain failed' in response['message']
|
||||
|
||||
|
||||
@@ -415,19 +566,17 @@ async def test_retrain_model_data_error(mlflow):
|
||||
|
||||
@mark.asyncio
|
||||
async def test_retrain_model_missing_target(mlflow):
|
||||
raw_data = MagicMock(columns=['variable', 'timestamp', 'value'])
|
||||
raw_data.__getitem__.return_value.max.return_value = 'ts'
|
||||
ts = pd.Timestamp('2020-01-01', tz='UTC')
|
||||
raw_data = pd.DataFrame(
|
||||
{
|
||||
'variable': ['f1', 'f2'],
|
||||
'timestamp': [ts, ts],
|
||||
'value': [1.0, 2.0],
|
||||
}
|
||||
)
|
||||
payload = AsyncMock()
|
||||
payload.retrieve = AsyncMock(return_value=raw_data)
|
||||
|
||||
pivoted = MagicMock()
|
||||
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 = await mlflow.retrain_model(
|
||||
{
|
||||
**metadata,
|
||||
@@ -561,6 +710,24 @@ async def test_get_reference_data_not_found(mlflow):
|
||||
assert result is None
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_get_reference_data_missing_csv_file_returns_none(mlflow):
|
||||
input_data = {
|
||||
**metadata,
|
||||
'model_name': 'test_model',
|
||||
}
|
||||
mv = MagicMock(run_id='run1')
|
||||
mlflow.mlflow_repository._client.get_model_version_by_alias.return_value = mv
|
||||
|
||||
with patch('laborious.activities.mlflow.tempfile.mkdtemp', return_value='tmp'):
|
||||
with patch('laborious.activities.mlflow.rmtree'):
|
||||
with patch('laborious.activities.mlflow.Path') as mp:
|
||||
mp.return_value.rglob.return_value = []
|
||||
result = await mlflow.get_reference_data(input_data)
|
||||
|
||||
assert result is None
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_get_reference_data_exception(mlflow):
|
||||
input_data = {
|
||||
|
||||
Reference in New Issue
Block a user