Update MLFlow class to pass logger instance directly to MLFlowRepository - Modified the initialization of MLFlowRepository in the MLFlow class to pass the logger instance directly, improving logging capabilities and consistency across the application.
340 lines
13 KiB
Python
340 lines
13 KiB
Python
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(data.head(5).to_string(), metadata)
|
|
|
|
# Convert numpy.nan to None for model compatibility
|
|
data.replace(np.nan, None, inplace=True)
|
|
|
|
# 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']
|
|
|
|
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 = data.dropna()
|
|
data.columns.name = None
|
|
|
|
try:
|
|
retrain_output, experiment = self.model_monitoring_repository.retrain_model(
|
|
data=data,
|
|
model_name=model_name
|
|
)
|
|
|
|
return {
|
|
'status': retrain_output,
|
|
'timestamp': timestamp,
|
|
'experiment': experiment
|
|
}
|
|
except Exception as e:
|
|
trace = traceback.format_exc()
|
|
self.send_notification(
|
|
metadata=metadata,
|
|
notification_id='RETRAIN_MODEL_ERROR',
|
|
message=f'Error retraining model {model_name}: {e}',
|
|
block='retrain_model',
|
|
level=NotificationLevel.ERROR,
|
|
attachment_content=trace
|
|
)
|
|
self.error(trace, metadata=metadata)
|
|
raise e
|
|
|
|
@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']
|
|
model_id = input_data['model_id']
|
|
experiment = input_data['experiment']
|
|
timestamp = input_data['timestamp']
|
|
status = input_data['status']
|
|
|
|
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
|
|
)
|
|
|
|
report = DataFrame([response])
|
|
report['model_id'] = model_id
|
|
report['model_name'] = model_name
|
|
report['timestamp'] = timestamp
|
|
report['status'] = status
|
|
|
|
self.info(
|
|
f'Production model {model_name} updated successfully', metadata)
|
|
return report.to_dict()
|
|
|
|
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
|