This commit renames the 'laborious' package to 'model_manager' across the entire project. This includes renaming directories, modules, references in code, configuration files, and documentation to reflect the new package name. This change improves clarity and consistency within the project.
298 lines
9.2 KiB
Python
298 lines
9.2 KiB
Python
from datetime import datetime
|
|
from unittest.mock import ANY, MagicMock, patch
|
|
|
|
import numpy as np
|
|
from pandas import DataFrame, Timestamp
|
|
from pytest import fixture, mark, raises
|
|
from sientia_do.temporal.constants import DATETIME_FORMAT, DATETIME_FORMAT_WITH_TZ
|
|
from model_manager.activities.mlflow import MLFlow
|
|
from sientia_do.notifications.models import NotificationLevel
|
|
|
|
|
|
@patch("model_manager.activities.mlflow.MLFlowRepository")
|
|
def test___init__(mock_mlflow_repository):
|
|
mlflow = MLFlow(
|
|
mlflow_host="http://localhost",
|
|
mlflow_port=5000,
|
|
mlflow_username="admin",
|
|
mlflow_password="admin",
|
|
logger=MagicMock(),
|
|
notification_handler=MagicMock()
|
|
)
|
|
|
|
assert mlflow.mlflow_host == "http://localhost"
|
|
assert mlflow.mlflow_port == 5000
|
|
assert mlflow.mlflow_username == "admin"
|
|
assert mlflow.mlflow_password == "admin"
|
|
|
|
mock_mlflow_repository.assert_called_once_with(
|
|
"http://localhost:5000", "admin", "admin", ANY
|
|
)
|
|
|
|
|
|
@fixture
|
|
@patch("model_manager.activities.mlflow.MLFlowRepository")
|
|
def mlflow(mock_mlflow_repository):
|
|
mlflow = MLFlow(
|
|
mlflow_host="http://localhost:5000",
|
|
mlflow_port=5000,
|
|
mlflow_username="admin",
|
|
mlflow_password="admin",
|
|
logger=MagicMock(),
|
|
notification_handler=MagicMock()
|
|
)
|
|
|
|
mlflow.send_notification = MagicMock()
|
|
|
|
return mlflow
|
|
|
|
|
|
metadata = {
|
|
"metadata": {
|
|
"model_id": "test_model",
|
|
"model_name": "test_model",
|
|
"workflow_name": "test_workflow",
|
|
"schema_name": "test_schedule",
|
|
},
|
|
}
|
|
|
|
|
|
@mark.asyncio
|
|
@patch("model_manager.activities.mlflow.DataFrame")
|
|
@patch("model_manager.activities.mlflow.max")
|
|
async def test_request_transform_success(mock_max, mock_dataframe, mlflow):
|
|
mock_max.return_value = '2024-01-02'
|
|
# Mock input data
|
|
input_data = {
|
|
**metadata,
|
|
'data': [
|
|
{'timestamp': '2024-01-01', 'variable': 'var1',
|
|
'value': 1.0, 'created_at': '2024-01-01 12:00:00'},
|
|
{'timestamp': '2024-01-01', 'variable': 'var2',
|
|
'value': 2.0, 'created_at': '2024-01-01 12:00:00'},
|
|
{'timestamp': '2024-01-02', 'variable': 'var1',
|
|
'value': 3.0, 'created_at': '2024-01-02 12:00:00'},
|
|
{'timestamp': '2024-01-02', 'variable': 'var2',
|
|
'value': 4.0, 'created_at': '2024-01-02 12:00:00'},
|
|
{'timestamp': '2024-01-02', 'variable': 'var1',
|
|
'value': 1.0, 'created_at': '2024-01-01 12:00:00'},
|
|
{'timestamp': '2024-01-02', 'variable': 'var2',
|
|
'value': 1.0, 'created_at': '2024-01-01 12:00:00'}
|
|
],
|
|
'model_name': 'test_model',
|
|
'model_config': {}
|
|
}
|
|
|
|
# Mock the transform response
|
|
expected_response = {'prediction': [0.5, 0.6], 'timestamp': [
|
|
'2024-01-01', '2024-01-02']}
|
|
mlflow.model_monitoring_repository.transform.return_value = expected_response
|
|
|
|
mock_dataframe.return_value.sort_values.return_value = mock_dataframe.return_value
|
|
mock_dataframe.return_value.drop_duplicates.return_value = mock_dataframe.return_value
|
|
|
|
# Call the method
|
|
response_data = await mlflow.request_transform(input_data)
|
|
|
|
# Verify the data was correctly transformed
|
|
mock_dataframe.assert_called_once_with(input_data['data'])
|
|
mock_dataframe.return_value.pivot.assert_called_once_with(
|
|
index='timestamp', columns='variable', values='value'
|
|
)
|
|
mock_dataframe = mock_dataframe.return_value.pivot.return_value
|
|
mock_dataframe.fillna.assert_called_once_with(np.nan, inplace=True)
|
|
# mock_dataframe.reset_index.assert_called_once()
|
|
mock_dataframe.columns.name = None
|
|
|
|
# Verify the response
|
|
assert response_data == expected_response
|
|
|
|
# Verify the repository was called with correct arguments
|
|
mlflow.model_monitoring_repository.transform.assert_called_once_with(
|
|
'test_model', mock_dataframe, {}, metadata['metadata']
|
|
)
|
|
|
|
|
|
@mark.asyncio
|
|
@patch("model_manager.activities.mlflow.DataFrame")
|
|
@patch("model_manager.activities.mlflow.to_datetime")
|
|
@patch("model_manager.activities.mlflow.max")
|
|
async def test_request_predict(mock_max, mock_to_datetime, mock_dataframe, mlflow):
|
|
mock_max.return_value = '2024-01-02'
|
|
# Mock input data
|
|
input_data = {
|
|
**metadata,
|
|
'data': {
|
|
"variable": {
|
|
"2024-01-01": "var1",
|
|
"2024-01-02": "var2",
|
|
"2024-01-03": "var1",
|
|
"2024-01-04": "var2"
|
|
},
|
|
"value": {
|
|
"2024-01-01": 1.0,
|
|
"2024-01-02": 2.0,
|
|
"2024-01-03": 3.0,
|
|
"2024-01-04": 4.0
|
|
}
|
|
},
|
|
'model_name': 'test_model',
|
|
'model_config': {}
|
|
}
|
|
|
|
# Mock the predict response
|
|
expected_response = {'prediction': [0.5, 0.6]}
|
|
mlflow.model_monitoring_repository.predict.return_value = expected_response
|
|
|
|
# Call the method
|
|
response_data = await mlflow.request_predict(input_data)
|
|
|
|
mock_dataframe.assert_called_once_with(input_data['data'])
|
|
mock_dataframe.return_value.replace.assert_called_once_with(
|
|
np.nan, None, inplace=True
|
|
)
|
|
mock_dataframe.return_value.__setitem__.assert_any_call(
|
|
'timestamp', mock_to_datetime.return_value.dt.strftime.return_value
|
|
)
|
|
|
|
mock_to_datetime.assert_called_once_with(
|
|
mock_dataframe.return_value.__getitem__.return_value, format=DATETIME_FORMAT_WITH_TZ
|
|
)
|
|
mock_to_datetime.return_value.dt.strftime.assert_called_once_with(
|
|
DATETIME_FORMAT
|
|
)
|
|
|
|
# Verify the response
|
|
assert response_data == expected_response
|
|
|
|
# Verify the repository was called with correct arguments
|
|
mlflow.model_monitoring_repository.predict.assert_called_once_with(
|
|
'test_model', mock_dataframe.return_value, {}, metadata['metadata']
|
|
)
|
|
|
|
|
|
@mark.asyncio
|
|
async def test_retrain_model(mlflow):
|
|
data = {
|
|
"model_id": [4, 5, 6, 7],
|
|
"created_at": [1, 2, 3, 4],
|
|
"timestamp": [1, 1, 2, 2],
|
|
"variable": ["var1", "var2", "var1", "var2"],
|
|
"value": [1, 2, 3, 4]
|
|
}
|
|
|
|
mlflow.model_monitoring_repository.retrain_model.return_value = (
|
|
'Model retrained successfully', 'test')
|
|
|
|
response = await mlflow.retrain_model({
|
|
**metadata,
|
|
'data': data,
|
|
'model_name': 'test_model'
|
|
})
|
|
|
|
mlflow.model_monitoring_repository.retrain_model.assert_called_once()
|
|
|
|
assert response == {
|
|
"status": 'Model retrained successfully',
|
|
"timestamp": 2,
|
|
"experiment": 'test'
|
|
}
|
|
|
|
|
|
@mark.asyncio
|
|
async def test_retrain_model_error(mlflow):
|
|
mlflow.model_monitoring_repository.retrain_model.side_effect = Exception(
|
|
'Error retraining model'
|
|
)
|
|
|
|
data = {
|
|
"model_id": [4, 5, 6, 7],
|
|
"created_at": [1, 2, 3, 4],
|
|
"timestamp": [1, 1, 2, 2],
|
|
"variable": ["var1", "var2", "var1", "var2"],
|
|
"value": [1, 2, 3, 4]
|
|
}
|
|
|
|
try:
|
|
await mlflow.retrain_model({
|
|
**metadata,
|
|
'data': data,
|
|
'model_name': 'test_model'
|
|
})
|
|
except Exception as e:
|
|
assert str(e) == 'Error retraining model'
|
|
mlflow.send_notification.assert_called_once_with(
|
|
metadata=metadata['metadata'],
|
|
notification_id='RETRAIN_MODEL_ERROR',
|
|
message='Error retraining model test_model: Error retraining model',
|
|
block='retrain_model',
|
|
level=NotificationLevel.ERROR,
|
|
attachment_content=ANY
|
|
)
|
|
else:
|
|
assert False, "No exception raised"
|
|
|
|
|
|
@mark.asyncio
|
|
async def test_update_production_model(mlflow):
|
|
mlflow.model_monitoring_repository.update_production_model.return_value = (
|
|
{
|
|
"data1": 1,
|
|
"data2": 2
|
|
}
|
|
)
|
|
|
|
input_data = {
|
|
**metadata,
|
|
'model_name': 'test_model',
|
|
'model_id': 1,
|
|
'experiment': 'test',
|
|
'timestamp': 2,
|
|
'status': 'success'
|
|
}
|
|
|
|
response = await mlflow.update_production_model(input_data)
|
|
|
|
mlflow.model_monitoring_repository.update_production_model.assert_called_once_with(
|
|
experiment='test', model_name='test_model')
|
|
|
|
assert response == {
|
|
'data1': {0: 1},
|
|
'data2': {0: 2},
|
|
'model_id': {0: 1},
|
|
'model_name': {0: 'test_model'},
|
|
'timestamp': {0: 2},
|
|
'status': {0: 'success'}
|
|
}
|
|
|
|
|
|
@mark.asyncio
|
|
async def test_update_production_model_error(mlflow):
|
|
mlflow.model_monitoring_repository.update_production_model.side_effect = Exception(
|
|
'Error updating production model'
|
|
)
|
|
|
|
input_data = {
|
|
**metadata,
|
|
'model_name': 'test_model',
|
|
'model_id': 1,
|
|
'experiment': 'test',
|
|
'timestamp': 2,
|
|
'status': 'success'
|
|
}
|
|
|
|
try:
|
|
await mlflow.update_production_model(input_data)
|
|
except Exception as e:
|
|
assert str(e) == 'Error updating production model'
|
|
mlflow.send_notification.assert_called_once_with(
|
|
metadata=metadata['metadata'],
|
|
notification_id='UPDATE_PRODUCTION_MODEL_ERROR',
|
|
message='Error updating production model test_model: Error updating production model',
|
|
block='update_production_model',
|
|
level=NotificationLevel.ERROR,
|
|
attachment_content=ANY
|
|
)
|
|
else:
|
|
assert False, "No exception raised"
|