SIENTIAPDE-994
Update requirements.txt with new dependencies and refactor activity methods for improved functionality and error handling
This commit is contained in:
@@ -0,0 +1,65 @@
|
||||
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.model_repository import ModelMonitoringRepository
|
||||
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 = ModelMonitoringRepository(
|
||||
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]:
|
||||
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]:
|
||||
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
|
||||
|
||||
Reference in New Issue
Block a user