from temporalio import activity, workflow with workflow.unsafe.imports_passed_through(): from datetime import datetime from pandas import Timestamp, to_datetime from sientia_do.temporal.constants import DATETIME_FORMAT, DATETIME_FORMAT_WITH_TZ from sientia_do.temporal.activities.base import BaseActivity from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler from sientia_do.notifications.models import NotificationLevel from sientia_do.observability.logger import Logger from sientia_do.formatters import create_sample_dict from laborious.utils.repository.model_repository import MLFlowRepository from typing import Any import numpy as np from pandas import DataFrame import traceback class MLFlow(BaseActivity): """ MLFlow integration activities for model inference operations. This class provides activities for interacting with MLFlow models, including data transformation and prediction operations. It handles authentication, data preprocessing, and model management with configurable retention policies. The class implements comprehensive error handling and logging for all MLFlow operations, ensuring reliable model inference in production environments. Attributes: mlflow_host (str): MLFlow server hostname mlflow_port (int): MLFlow server port mlflow_username (str): MLFlow authentication username mlflow_password (str): MLFlow authentication password model_monitoring_repository (MLFlowRepository): Repository for MLFlow operations """ def __init__(self, mlflow_host: str, mlflow_port: int, mlflow_username: str, mlflow_password: str, logger: Logger, notification_handler: NotificationHandler): """ Initialize MLFlow activities with server configuration. Args: mlflow_host: MLFlow server hostname or IP address mlflow_port: MLFlow server port number mlflow_username: Username for MLFlow authentication mlflow_password: Password for MLFlow authentication logger: Logger instance for observability and debugging notification_handler: Notification handler for alerts and monitoring Raises: Exception: If MLFlowRepository initialization fails """ BaseActivity.__init__( self, logger, notification_handler, set_error_counter=True) 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, logger ) @activity.defn(name="request_transform") async def request_transform(self, input_data: dict[str, Any]) -> dict[str, Any]: """ Transform input data using MLFlow models. This activity processes input data through MLFlow model transformation, including data preprocessing, format conversion, and validation. It handles data deduplication, pivoting, and cleanup to ensure optimal model performance. The transformation process includes: 1. Data deduplication based on variable and timestamp 2. Data pivoting for model input format 3. Null value handling and cleanup 4. MLFlow model transformation request 5. Response validation and logging Args: input_data: Configuration and data for transformation Required keys: - metadata (dict): Workflow execution metadata - data (dict): Input data for transformation - model_name (str): Name of the MLFlow model to use - model_retention (int): Model retention period in minutes Returns: dict: Transformed data from MLFlow model Raises: Exception: If transformation fails or MLFlow model is unavailable """ metadata = input_data['metadata'] self.info('Transforming data...', metadata) data = DataFrame(input_data['data']) model_name = input_data['model_name'] model_config = input_data.get('model_config', {}) self.debug("Raw input data:", metadata) self.debug(data.head(5).to_string(), metadata) # Sort by created_at in descending order and keep first occurrence of each variable/timestamp pair data = data.sort_values('created_at', ascending=False).drop_duplicates( subset=['variable', 'timestamp'], keep='first' ) # Pivot data for model input format data = data.pivot( index='timestamp', columns='variable', values='value') data.fillna(np.nan, inplace=True) # data.reset_index(inplace=True) data.columns.name = None self.debug("Processed input data:", metadata) self.debug(data.head(5).to_string(), metadata) # Request transformation from MLFlow model response_data = self.model_monitoring_repository.transform( model_name, data, model_config, metadata ) self.debug( f"Transform raw response data: \n {create_sample_dict(response_data, max_items=5, max_depth=5)}", metadata) self.debug( f"Transform response data: \n {create_sample_dict(response_data, max_items=5, max_depth=5)}", metadata) self.info("Data transformed successfully", metadata) return response_data @activity.defn(name="request_predict") async def request_predict(self, input_data: dict[str, Any]) -> dict[str, Any]: """ Execute predictions using MLFlow models. This activity performs ML model inference using MLFlow models with the transformed data. It handles data format conversion, null value processing, and model prediction requests with comprehensive error handling. The prediction process includes: 1. Data format validation and cleanup 2. Null value handling for model compatibility 3. MLFlow model prediction request 4. Response validation and logging 5. Performance monitoring and metrics Args: input_data: Configuration and data for prediction Required keys: - metadata (dict): Workflow execution metadata - data (dict): Transformed data for prediction - model_name (str): Name of the MLFlow model to use - model_retention (int): Model retention period in minutes Returns: dict: Prediction results from MLFlow model Raises: Exception: If prediction fails or MLFlow model is unavailable """ metadata = input_data['metadata'] self.info('Predicting data...', metadata) data = DataFrame(input_data['data']) model_name = input_data['model_name'] model_config = input_data.get('model_config', {}) self.debug(f"Input data for: \n {data.head(5).to_string()}", metadata) # Convert numpy.nan to None for model compatibility data.replace(np.nan, None, inplace=True) data['timestamp'] = data.index data['timestamp'] = to_datetime( data['timestamp'], format=DATETIME_FORMAT_WITH_TZ).dt.strftime(DATETIME_FORMAT) # Request prediction from MLFlow model response_data = self.model_monitoring_repository.predict( model_name, data, model_config, metadata ) self.debug( f"Prediction response data: \n {create_sample_dict(response_data, max_items=5, max_depth=5)}", metadata) self.info("Data predicted successfully", metadata) return response_data @activity.defn(name="retrain_model") async def retrain_model(self, input_data: dict[str, Any]) -> dict[str, Any]: """ Retrain MLFlow models with updated training data. This activity orchestrates the complete model retraining process, including data preparation, model retraining execution, and result validation. It handles data preprocessing, column cleanup, and comprehensive error handling for production model management. The retraining process includes: 1. Data timestamp extraction and validation 2. Column cleanup and data preparation 3. Data pivoting for model input format 4. MLFlow model retraining execution 5. Result validation and error handling Args: input_data (dict): Input data containing: - metadata (dict): Workflow execution metadata - data (dict[str, Any]): Training data for model retraining - model_name (str): Name of the MLFlow model to retrain Returns: dict: Retraining results containing: - status (str): Retraining operation status - timestamp (str): Timestamp of the retraining operation - experiment (str): MLFlow experiment identifier Raises: Exception: If retraining fails or encounters critical errors """ metadata = input_data['metadata'] data = DataFrame(input_data['data']) model_name = input_data['model_name'] model_config = input_data.get('model_config', {}) self.info(f'Retraining model {model_name}...', metadata) timestamp = data['timestamp'].max() self.debug(f'Timestamp: {timestamp}', metadata) data.drop(columns=['model_id'], inplace=True, errors='ignore') data.drop(columns=['created_at'], inplace=True, errors='ignore') data = data.pivot(index='timestamp', columns='variable', values='value') data.sort_index(inplace=True) data.reset_index(inplace=True) data['timestamp'] = to_datetime( data['timestamp'], format=DATETIME_FORMAT_WITH_TZ).dt.strftime(DATETIME_FORMAT) data['timestamp'] = to_datetime( data['timestamp'], format=DATETIME_FORMAT) # data = data.dropna() data.columns.name = None retrain_output = self.model_monitoring_repository.retrain_model( data=data, model_name=model_name, model_config=model_config, metadata=metadata ) if not retrain_output['success']: trace = retrain_output['traceback'] self.send_notification( metadata=metadata, notification_id='RETRAIN_MODEL_ERROR', message=f"Error retraining model {model_name}: {retrain_output['message']}", block='retrain_model', level=NotificationLevel.ERROR, attachment_content=trace ) self.error(trace, metadata=metadata) return { **retrain_output, 'timestamp': timestamp } @activity.defn(name="update_production_model") async def update_production_model(self, input_data: dict[str, Any]) -> dict[Any, Any]: """ Update production model with newly trained model version. This activity manages the critical process of updating production models with newly trained versions. It handles model deployment, status tracking, and comprehensive reporting for operational visibility and audit trails. The update process includes: 1. Production model update execution 2. Status and metadata tracking 3. Comprehensive reporting and logging 4. Error handling and notification 5. Audit trail maintenance Args: input_data (dict): Input data containing: - metadata (dict): Workflow execution metadata - model_name (str): Name of the MLFlow model to update - experiment (str): MLFlow experiment identifier - model_id (str): Unique identifier for the model version - timestamp (str): Timestamp of the update operation - status (str): Current status of the model update Returns: dict[Any, Any]: Comprehensive update report containing: - model_id (str): Model version identifier - model_name (str): Name of the updated model - timestamp (str): Update operation timestamp - status (str): Update operation status - Additional MLFlow response metadata Raises: Exception: If production model update fails """ metadata = input_data['metadata'] model_name = input_data['model_name'] experiment = input_data['experiment'] self.info( f'Updating production model {model_name} from experiment {experiment}...', metadata) try: response = self.model_monitoring_repository.update_production_model( experiment=experiment, model_name=model_name ) self.info( f'Production model {model_name} updated successfully', metadata) return response except Exception as e: trace = traceback.format_exc() self.send_notification( metadata=metadata, notification_id='UPDATE_PRODUCTION_MODEL_ERROR', message=f'Error updating production model {model_name}: {e}', block='update_production_model', level=NotificationLevel.ERROR, attachment_content=trace ) self.error(trace, metadata=metadata) raise e