451 lines
18 KiB
Python
451 lines
18 KiB
Python
from temporalio import activity, workflow
|
|
|
|
with workflow.unsafe.imports_passed_through():
|
|
import traceback
|
|
from typing import Any
|
|
|
|
import numpy as np
|
|
from pandas import DataFrame, to_datetime
|
|
from sientia_do.formatters import create_sample_dict
|
|
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.temporal.activities.base import BaseActivity
|
|
from sientia_do.temporal.constants import DATETIME_FORMAT, DATETIME_FORMAT_WITH_TZ
|
|
|
|
from model_manager.utils.repository.model_repository import MLFlowRepository
|
|
|
|
|
|
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']
|
|
|
|
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() # type: ignore[no-any-return]
|
|
|
|
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
|
|
|
|
@activity.defn(name='save_model')
|
|
async def save_model(self, input_data: dict[str, Any]) -> dict[str, Any]:
|
|
"""
|
|
Save a trained ML model and its artifacts to MLflow with comprehensive error handling.
|
|
|
|
This activity orchestrates the complete model saving pipeline:
|
|
1. Generates the next run name for the experiment
|
|
2. Creates and organizes artifacts (reports, data files)
|
|
3. Logs model, parameters, metrics, and artifacts to MLflow
|
|
4. Returns success/failure status with results or error message
|
|
|
|
The activity does NOT raise exceptions on failure - it catches all errors,
|
|
sends notifications, and returns a failure status. This allows the workflow
|
|
to handle the error gracefully and update the database accordingly.
|
|
|
|
Args:
|
|
input_data: Configuration for model saving operation
|
|
Required keys:
|
|
- metadata (dict): Workflow execution metadata
|
|
- train_result (TrainModelResult): Training result with model and metrics
|
|
|
|
Returns:
|
|
dict: Save result with the following structure:
|
|
{
|
|
'success': bool, # True if saving succeeded, False otherwise
|
|
'result': TrainModelResult | None, # Updated result if success=True
|
|
'error_message': str | None # Error message if success=False
|
|
}
|
|
|
|
Example:
|
|
# Successful save
|
|
result = await save_model({
|
|
'metadata': {'workflow_id': 'save-123', 'experiment_run_id': 456},
|
|
'train_result': TrainModelResult(...)
|
|
})
|
|
# Returns: {'success': True, 'result': TrainModelResult(...), 'error_message': None}
|
|
|
|
# Failed save
|
|
# Returns: {'success': False, 'result': None, 'error_message': 'Error details...'}
|
|
"""
|
|
metadata = input_data.get('metadata', {})
|
|
train_result = input_data['train_result']
|
|
|
|
try:
|
|
experiment_name = train_result.params.experiment_name
|
|
|
|
self.info(
|
|
f'Starting model save for experiment: {experiment_name}',
|
|
metadata,
|
|
)
|
|
|
|
# Step 1: Generate next run name
|
|
self.info('Generating run name', metadata)
|
|
train_result.run_name = self.model_monitoring_repository.get_next_run_name(
|
|
experiment_name
|
|
)
|
|
self.info(f'Generated run name: {train_result.run_name}', metadata)
|
|
|
|
# Step 2: Generate artifacts (reports, CSV files)
|
|
self.info('Generating artifacts', metadata)
|
|
train_result = self.model_monitoring_repository.generate_artifacts(train_result)
|
|
self.info('Artifacts generated successfully', metadata)
|
|
|
|
# Step 3: Save run to MLflow
|
|
self.info('Saving run to MLflow', metadata)
|
|
self.model_monitoring_repository.save_run(train_result)
|
|
|
|
self.info(
|
|
f'Model saved successfully - Run: {train_result.run_name}, '
|
|
f'Experiment: {experiment_name}',
|
|
metadata,
|
|
)
|
|
|
|
return {
|
|
'success': True,
|
|
'result': train_result,
|
|
'error_message': None,
|
|
}
|
|
|
|
except Exception as e: # noqa: BLE001
|
|
error_msg = f'Error saving model - Experiment: {train_result.params.experiment_name if train_result and train_result.params else "unknown"}, Error: {str(e)}'
|
|
trace = traceback.format_exc()
|
|
|
|
# Send notification (MongoDB)
|
|
self.send_notification(
|
|
metadata=metadata,
|
|
notification_id='SAVE_MODEL_ERROR',
|
|
message=error_msg,
|
|
block='save_model',
|
|
level=NotificationLevel.ERROR,
|
|
attachment_content=trace,
|
|
)
|
|
|
|
# Log error with metadata
|
|
self.error(trace, metadata=metadata)
|
|
|
|
# Return failure result (do NOT raise exception)
|
|
# This allows workflow to update database with error status
|
|
return {
|
|
'success': False,
|
|
'result': None,
|
|
'error_message': str(e),
|
|
}
|