SIENTIAPDE-1222 Enhance MLFlow logging with sample dictionary for response data - Introduced a new method to create a sample dictionary for debugging, allowing for better visualization of nested data structures in logs. - Updated debug logging to utilize the new sampling method for raw and transformed response data, improving clarity and reducing output size. - Adjusted logging for processed input data to display only the first few rows, enhancing readability.
443 lines
17 KiB
Python
443 lines
17 KiB
Python
from temporalio import activity, workflow
|
|
|
|
|
|
with workflow.unsafe.imports_passed_through():
|
|
from datetime import datetime
|
|
import json
|
|
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 laborious.utils.repository.model_repository import MLFlowRepository
|
|
from typing import Any
|
|
import numpy as np
|
|
from pandas import DataFrame
|
|
import traceback
|
|
|
|
|
|
def create_sample_dict(data: dict, max_items: int = 3, max_depth: int = 2) -> dict:
|
|
"""
|
|
Create a sample of a dictionary for debugging purposes.
|
|
|
|
Args:
|
|
data: Dictionary to sample
|
|
max_items: Maximum number of items to show per level
|
|
max_depth: Maximum depth to traverse nested structures
|
|
|
|
Returns:
|
|
Dictionary with sampled content
|
|
"""
|
|
if max_depth <= 0:
|
|
return {"...": "max_depth_reached"}
|
|
|
|
sample = {}
|
|
items = list(data.items())[:max_items]
|
|
|
|
for key, value in items:
|
|
if isinstance(value, dict):
|
|
sample[key] = create_sample_dict(value, max_items, max_depth - 1)
|
|
elif isinstance(value, list):
|
|
sample[key] = value[:max_items] if len(
|
|
value) > max_items else value
|
|
else:
|
|
sample[key] = value
|
|
|
|
if len(data) > max_items:
|
|
sample["..."] = f"({len(data) - max_items} more items)"
|
|
|
|
return sample
|
|
|
|
|
|
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.base_logger
|
|
)
|
|
|
|
def detect_and_parse_datetime_index(self, data: DataFrame, metadata: dict) -> DataFrame:
|
|
"""
|
|
Detect and parse datetime index from data. index must be a timestamp like column.
|
|
This function must detect the timestamp type (pandas Timestamp or datetime) and convert it to DATETIME_FORMAT_WITH_TZ.
|
|
If the index is a string, must be in format DATETIME_FORMAT_WITH_TZ.
|
|
If another type or format, must raise an error.
|
|
"""
|
|
index = data.index
|
|
|
|
# Get type of first element of index
|
|
index_type = type(index[0])
|
|
|
|
self.info(f"Index type: {index_type}", metadata)
|
|
|
|
message = f"Index must be all timestamp like column. Valid formats are: pandas Timestamp, datetime, string in format {DATETIME_FORMAT_WITH_TZ}"
|
|
|
|
# Check if all in index are of the same type
|
|
if not all(isinstance(i, index_type) for i in index):
|
|
raise ValueError(
|
|
f"{message}")
|
|
|
|
# Check type and converts to DATETIME_FORMAT_WITH_TZ
|
|
if index_type == str:
|
|
# Validate format of string and return error if not valid
|
|
try:
|
|
to_datetime(data.index)
|
|
except ValueError:
|
|
raise ValueError(
|
|
f"{message}")
|
|
|
|
elif index_type == datetime or index_type == Timestamp:
|
|
data.index = data.index.strftime(DATETIME_FORMAT_WITH_TZ)
|
|
else:
|
|
raise ValueError(
|
|
f"{message}")
|
|
|
|
return data
|
|
|
|
@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
|
|
)
|
|
|
|
self.debug("Transform raw response data:", metadata)
|
|
self.debug(create_sample_dict(
|
|
response_data), metadata)
|
|
|
|
if response_data['success']:
|
|
|
|
response_dataframe = DataFrame(response_data['content'])
|
|
if len(response_dataframe) == 0:
|
|
return response_data
|
|
try:
|
|
response_dataframe = self.detect_and_parse_datetime_index(
|
|
response_dataframe, metadata)
|
|
response_dataframe['timestamp'] = to_datetime(
|
|
response_dataframe.index, format=DATETIME_FORMAT_WITH_TZ)
|
|
response_dataframe['timestamp'] = response_dataframe['timestamp'].dt.strftime(
|
|
DATETIME_FORMAT)
|
|
except ValueError as e:
|
|
trace = traceback.format_exc()
|
|
self.send_notification(
|
|
metadata=metadata,
|
|
notification_id='TRANSFORM_DATA_INDEX_ERROR',
|
|
message=f'Error parsing trasnformed data index: {e}',
|
|
block='transform',
|
|
level=NotificationLevel.ERROR,
|
|
attachment_content=trace
|
|
)
|
|
self.error(trace, metadata=metadata)
|
|
raise e
|
|
|
|
response_dataframe.to_csv('response_data.csv')
|
|
|
|
response_data['content'] = response_dataframe.to_dict()
|
|
|
|
self.debug("Transform response data:", metadata)
|
|
self.debug(create_sample_dict(
|
|
response_data), 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
|
|
)
|
|
|
|
self.debug("Prediction response data:", metadata)
|
|
self.debug(create_sample_dict(
|
|
response_data), 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
|