SIENTIAPDE-1222
Refactor datetime index handling in MLFlow and MLFlowRepository - Moved the detect_and_parse_datetime_index method from MLFlow to MLFlowRepository for better organization and reusability. - Updated the method to include enhanced logging and error handling for invalid datetime formats. - Adjusted the transform method in MLFlowRepository to utilize the new datetime index parsing logic. - Added unit tests for both valid and invalid datetime index cases to ensure robustness.
This commit is contained in:
@@ -63,44 +63,6 @@ class MLFlow(BaseActivity):
|
||||
f"{mlflow_host}:{mlflow_port}", mlflow_username, mlflow_password, logger.base_logger
|
||||
)
|
||||
|
||||
def detect_and_parse_datetime_index(self, data: DataFrame, metadata: dict) -> DataFrame:
|
||||
"""
|
||||
Detect and parse datetime index from data. index must be a timestamp like column.
|
||||
This function must detect the timestamp type (pandas Timestamp or datetime) and convert it to DATETIME_FORMAT_WITH_TZ.
|
||||
If the index is a string, must be in format DATETIME_FORMAT_WITH_TZ.
|
||||
If another type or format, must raise an error.
|
||||
"""
|
||||
index = data.index
|
||||
|
||||
# Get type of first element of index
|
||||
index_type = type(index[0])
|
||||
|
||||
self.info(f"Index type: {index_type}", metadata)
|
||||
|
||||
message = f"Index must be all timestamp like column. Valid formats are: pandas Timestamp, datetime, string in format {DATETIME_FORMAT_WITH_TZ}"
|
||||
|
||||
# Check if all in index are of the same type
|
||||
if not all(isinstance(i, index_type) for i in index):
|
||||
raise ValueError(
|
||||
f"{message}")
|
||||
|
||||
# Check type and converts to DATETIME_FORMAT_WITH_TZ
|
||||
if index_type == str:
|
||||
# Validate format of string and return error if not valid
|
||||
try:
|
||||
to_datetime(data.index)
|
||||
except ValueError:
|
||||
raise ValueError(
|
||||
f"{message}")
|
||||
|
||||
elif index_type == datetime or index_type == Timestamp:
|
||||
data.index = data.index.strftime(DATETIME_FORMAT_WITH_TZ)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"{message}")
|
||||
|
||||
return data
|
||||
|
||||
@activity.defn(name="request_transform")
|
||||
async def request_transform(self, input_data: dict[str, Any]) -> dict[str, Any]:
|
||||
"""
|
||||
@@ -164,35 +126,6 @@ class MLFlow(BaseActivity):
|
||||
self.debug(
|
||||
f"Transform raw response data: \n {create_sample_dict(response_data, max_items=5, max_depth=5)}", metadata)
|
||||
|
||||
if response_data['success']:
|
||||
|
||||
response_dataframe = DataFrame(response_data['content'])
|
||||
if len(response_dataframe) == 0:
|
||||
return response_data
|
||||
try:
|
||||
response_dataframe = self.detect_and_parse_datetime_index(
|
||||
response_dataframe, metadata)
|
||||
response_dataframe['timestamp'] = to_datetime(
|
||||
response_dataframe.index, format=DATETIME_FORMAT_WITH_TZ)
|
||||
response_dataframe['timestamp'] = response_dataframe['timestamp'].dt.strftime(
|
||||
DATETIME_FORMAT)
|
||||
except ValueError as e:
|
||||
trace = traceback.format_exc()
|
||||
self.send_notification(
|
||||
metadata=metadata,
|
||||
notification_id='TRANSFORM_DATA_INDEX_ERROR',
|
||||
message=f'Error parsing trasnformed data index: {e}',
|
||||
block='transform',
|
||||
level=NotificationLevel.ERROR,
|
||||
attachment_content=trace
|
||||
)
|
||||
self.error(trace, metadata=metadata)
|
||||
raise e
|
||||
|
||||
response_dataframe.to_csv('response_data.csv')
|
||||
|
||||
response_data['content'] = response_dataframe.to_dict()
|
||||
|
||||
self.debug(
|
||||
f"Transform response data: \n {create_sample_dict(response_data, max_items=5, max_depth=5)}", metadata)
|
||||
|
||||
|
||||
@@ -16,17 +16,57 @@ import pandas as pd
|
||||
import mlflow
|
||||
from os import makedirs, path, remove
|
||||
from sientia.ModelServing import ModelServing
|
||||
from sientia_do.temporal.constants import DATETIME_FORMAT_WITH_TZ
|
||||
from sientia_do.observability.logger import Logger
|
||||
|
||||
|
||||
class MLFlowRepository():
|
||||
def __init__(self, host, username, password, logger):
|
||||
def __init__(self, host, username, password, logger: Logger):
|
||||
|
||||
self.model_serving = ModelServing(tracking_uri=host,
|
||||
username=username, password=password,
|
||||
logger=logger)
|
||||
self.logger = logger
|
||||
|
||||
def transform(self, model_name: str, data: pd.DataFrame, model_config: dict) -> dict:
|
||||
def detect_and_parse_datetime_index(self, data: pd.DataFrame, metadata: dict) -> pd.DataFrame:
|
||||
"""
|
||||
Detect and parse datetime index from data. index must be a timestamp like column.
|
||||
This function must detect the timestamp type (pandas Timestamp or datetime) and convert it to DATETIME_FORMAT_WITH_TZ.
|
||||
If the index is a string, must be in format DATETIME_FORMAT_WITH_TZ.
|
||||
If another type or format, must raise an error.
|
||||
"""
|
||||
index = data.index
|
||||
|
||||
# Get type of first element of index
|
||||
index_type = type(index[0])
|
||||
|
||||
self.logger.custom_info(f"Index type: {index_type}", metadata)
|
||||
|
||||
message = f"Index must be all timestamp like column. Valid formats are: pandas Timestamp, datetime, string in format {DATETIME_FORMAT_WITH_TZ}"
|
||||
|
||||
# Check if all in index are of the same type
|
||||
if not all(isinstance(i, index_type) for i in index):
|
||||
raise ValueError(
|
||||
f"{message}")
|
||||
|
||||
# Check type and converts to DATETIME_FORMAT_WITH_TZ
|
||||
if index_type == str:
|
||||
# Validate format of string and return error if not valid
|
||||
try:
|
||||
pd.to_datetime(data.index, format=DATETIME_FORMAT_WITH_TZ)
|
||||
except ValueError:
|
||||
raise ValueError(
|
||||
f"{message}")
|
||||
|
||||
elif index_type == datetime or index_type == pd.Timestamp:
|
||||
data.index = data.index.strftime(DATETIME_FORMAT_WITH_TZ)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"{message}")
|
||||
|
||||
return data
|
||||
|
||||
def transform(self, model_name: str, data: pd.DataFrame, model_config: dict, metadata: dict) -> dict:
|
||||
"""
|
||||
Transform data using a model.
|
||||
|
||||
@@ -40,6 +80,9 @@ class MLFlowRepository():
|
||||
"""
|
||||
|
||||
try:
|
||||
self.logger.custom_debug(
|
||||
f"Data received for model transformation: {data.to_csv()}", metadata)
|
||||
|
||||
model_retention = model_config.get('retention_minutes', 0)
|
||||
flavor = model_config.get('transform_flavor', 'sklearn')
|
||||
compressed = model_config.get('is_compressed', False)
|
||||
@@ -47,12 +90,20 @@ class MLFlowRepository():
|
||||
transform_keyword = model_config.get(
|
||||
'transform_function_keyword', 'predict')
|
||||
|
||||
transformed_data = self.model_serving.get_cached_transform(
|
||||
model_name, data, model_retention, flavor,
|
||||
compressed, retention_target, transform_keyword
|
||||
)
|
||||
|
||||
self.logger.custom_debug(
|
||||
f"Data received from model transformation: {transformed_data.to_csv()}", metadata)
|
||||
|
||||
transformed_data = self.detect_and_parse_datetime_index(
|
||||
transformed_data, metadata)
|
||||
|
||||
return {
|
||||
'success': True,
|
||||
'content': self.model_serving.get_cached_transform(
|
||||
model_name, data, model_retention, flavor,
|
||||
compressed, retention_target, transform_keyword
|
||||
).to_dict()
|
||||
'content': transformed_data.to_dict()
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
@@ -64,7 +115,7 @@ class MLFlowRepository():
|
||||
}
|
||||
}
|
||||
|
||||
def predict(self, model_name: str, data: pd.DataFrame, model_config: dict) -> dict:
|
||||
def predict(self, model_name: str, data: pd.DataFrame, model_config: dict, metadata: dict) -> dict:
|
||||
"""
|
||||
Predict data using a model.
|
||||
|
||||
@@ -85,8 +136,8 @@ class MLFlowRepository():
|
||||
input_index = data.index
|
||||
start_time = datetime.now()
|
||||
|
||||
self.logger.debug(
|
||||
f"Data received for model prediction: {data.to_string()}")
|
||||
self.logger.custom_debug(
|
||||
f"Data received for model prediction: {data.to_csv()}", metadata)
|
||||
data = self.model_serving.get_cached_predict(
|
||||
model_name, data, model_retention, flavor,
|
||||
compressed, retention_target
|
||||
@@ -94,8 +145,8 @@ class MLFlowRepository():
|
||||
|
||||
end_time = datetime.now()
|
||||
data = pd.DataFrame(data, columns=['prediction'])
|
||||
self.logger.debug(
|
||||
f"Data received from model prediction: {data.to_string()}")
|
||||
self.logger.custom_debug(
|
||||
f"Data received from model prediction: {data.to_csv()}", metadata)
|
||||
data.index = input_index
|
||||
data['response_time'] = (end_time - start_time).total_seconds()
|
||||
|
||||
|
||||
@@ -1,8 +1,9 @@
|
||||
from datetime import datetime
|
||||
from unittest.mock import ANY, MagicMock, patch
|
||||
|
||||
import numpy as np
|
||||
from pandas import DataFrame
|
||||
from pytest import fixture, mark
|
||||
from pandas import DataFrame, Timestamp
|
||||
from pytest import fixture, mark, raises
|
||||
from laborious.activities.mlflow import MLFlow
|
||||
from sientia_do.notifications.models import NotificationLevel
|
||||
|
||||
@@ -58,7 +59,7 @@ metadata = {
|
||||
@mark.asyncio
|
||||
@patch("laborious.activities.mlflow.DataFrame")
|
||||
@patch("laborious.activities.mlflow.max")
|
||||
async def test_request_transform(mock_max, mock_dataframe, mlflow):
|
||||
async def test_request_transform_success(mock_max, mock_dataframe, mlflow):
|
||||
mock_max.return_value = '2024-01-02'
|
||||
# Mock input data
|
||||
input_data = {
|
||||
|
||||
@@ -2,7 +2,8 @@ from unittest.mock import ANY, MagicMock, call, patch
|
||||
import numpy as np
|
||||
from pandas import DataFrame
|
||||
import pytest
|
||||
from laborious.utils.repository import model_repository
|
||||
from datetime import datetime, timezone
|
||||
from pandas import Timestamp
|
||||
from laborious.utils.repository.model_repository import MLFlowRepository
|
||||
|
||||
|
||||
@@ -22,18 +23,113 @@ def mlflow_repository():
|
||||
return repo
|
||||
|
||||
|
||||
metadata = {
|
||||
"metadata": {
|
||||
"model_id": "test_model",
|
||||
"model_name": "test_model",
|
||||
"workflow_name": "test_workflow",
|
||||
"schema_name": "test_schedule",
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
invalid_cases = [
|
||||
(
|
||||
{
|
||||
'value': {
|
||||
'2024-01-01 12:00:00': 1,
|
||||
2024: 2
|
||||
}
|
||||
}
|
||||
),
|
||||
(
|
||||
{
|
||||
'value': {
|
||||
'2024-01-01': 1,
|
||||
'2024-01-02': 2
|
||||
}
|
||||
}
|
||||
),
|
||||
(
|
||||
{
|
||||
'value': {
|
||||
1: 1,
|
||||
2: 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=timezone.utc): 1,
|
||||
datetime(2025, 1, 2, 12, 0, 0, tzinfo=timezone.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=timezone.utc): 1,
|
||||
Timestamp(2026, 1, 2, 12, 0, 0, tzinfo=timezone.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 = 'data'
|
||||
model_name = 'model'
|
||||
|
||||
output = mlflow_repository.transform(model_name, data, 1)
|
||||
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, 1)
|
||||
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.model_serving.get_cached_transform.return_value.to_dict.return_value
|
||||
'content': mlflow_repository.detect_and_parse_datetime_index.return_value.to_dict.return_value
|
||||
}
|
||||
|
||||
|
||||
@@ -44,10 +140,11 @@ def test_transform_error(mlflow_repository):
|
||||
mlflow_repository.model_serving.get_cached_transform.side_effect = Exception(
|
||||
'error')
|
||||
|
||||
output = mlflow_repository.transform(model_name, data, 1)
|
||||
output = mlflow_repository.transform(
|
||||
model_name, data, {}, metadata['metadata'])
|
||||
|
||||
mlflow_repository.model_serving.get_cached_transform.assert_called_once_with(
|
||||
model_name, data, 1)
|
||||
model_name, data, 0, 'sklearn', False, 'model', 'predict')
|
||||
|
||||
assert output == {
|
||||
'success': False,
|
||||
|
||||
Reference in New Issue
Block a user