Refactor activity methods and update requirements.txt to enhance functionality and remove deprecated filters. Added detailed docstrings for clarity and improved error handling in data processing workflows.
86 lines
3.2 KiB
Python
86 lines
3.2 KiB
Python
import numpy as np
|
|
from pandas import DataFrame
|
|
from temporalio import activity, workflow
|
|
|
|
|
|
with workflow.unsafe.imports_passed_through():
|
|
from laborious.activities.base import BaseActivity
|
|
from typing import Any
|
|
from logging import Logger
|
|
from sientia_do.notifications.handlers import NotificationHandler
|
|
from laborious.utils.repository.model_repository import MLFlowRepository
|
|
from sientia_do.notifications.models import NotificationLevel
|
|
|
|
|
|
class MLFlow(BaseActivity):
|
|
def __init__(self, mlflow_host: str, mlflow_port: int, mlflow_username: str,
|
|
mlflow_password: str, logger: Logger, notification_handler: NotificationHandler):
|
|
super().__init__(logger, notification_handler)
|
|
self.mlflow_host = mlflow_host
|
|
self.mlflow_port = mlflow_port
|
|
self.mlflow_username = mlflow_username
|
|
self.mlflow_password = mlflow_password
|
|
|
|
self.model_monitoring_repository = MLFlowRepository(
|
|
f"{mlflow_host}:{mlflow_port}", mlflow_username, mlflow_password
|
|
)
|
|
|
|
@activity.defn(name="request_transform")
|
|
async def request_transform(self, input_data: dict[str, Any]) -> tuple[dict[str, Any], str]:
|
|
"""
|
|
Access MLFlow model to get the transformed data.
|
|
Args:
|
|
input_data (dict): The input data. Contains:
|
|
data (dict[str, Any]): The data to transform.
|
|
model_name (str): The name of the model.
|
|
model_retention (int): The retention of the model.
|
|
Returns:
|
|
tuple[dict[str, Any], str]: The transformed data and the latest timestamp of the data.
|
|
"""
|
|
self.logger.info('Transforming data...')
|
|
data = DataFrame(input_data['data'])
|
|
model_name = input_data['model_name']
|
|
model_retention = input_data['model_retention']
|
|
|
|
self.logger.debug(data)
|
|
|
|
data = data.pivot(
|
|
index='timestamp', columns='variable',
|
|
values='value')
|
|
data.fillna(np.nan, inplace=True)
|
|
data.reset_index(inplace=True)
|
|
data.columns.name = None
|
|
|
|
response_data = self.model_monitoring_repository.transform(
|
|
model_name, data, model_retention)
|
|
|
|
timestamp = max(data['timestamp'].values.tolist())
|
|
|
|
return response_data, timestamp
|
|
|
|
@activity.defn(name="request_predict")
|
|
async def request_predict(self, input_data: dict[str, Any]) -> tuple[dict[str, Any], str]:
|
|
"""
|
|
Access MLFlow model to get the predicted data.
|
|
Args:
|
|
input_data (dict): The input data. Contains:
|
|
data (dict[str, Any]): The data to predict.
|
|
model_name (str): The name of the model.
|
|
model_retention (int): The retention of the model.
|
|
Returns:
|
|
tuple[dict[str, Any], str]: The predicted data and the latest timestamp of the data.
|
|
"""
|
|
self.logger.info('Predicting data...')
|
|
data = DataFrame(input_data['data'])
|
|
model_name = input_data['model_name']
|
|
model_retention = input_data['model_retention']
|
|
|
|
self.logger.debug(data)
|
|
|
|
data.replace(np.nan, None, inplace=True)
|
|
|
|
response_data = self.model_monitoring_repository.predict(
|
|
model_name, data, model_retention)
|
|
|
|
return response_data
|