Files
sientia-dataops-laborious_t…/laborious/activities/mlflow.py
vitor-aignosi fa88572126 SIENTIAPDE-1712
SIENTIAPDE-1712 Enhance logging in Gates, MLFlow, and Storage classes by integrating logger parameter for improved traceability. This update allows for better monitoring of operations and data handling across these components.
2026-03-30 09:24:52 -03:00

528 lines
20 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 to_datetime
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.observability.metrics_controller import MetricsController
from sientia_do.repository.minio_repository import MinioRepository
from sientia_do.temporal.constants import (
DATETIME_FORMAT,
DATETIME_FORMAT_MS_WITH_TZ,
DATETIME_FORMAT_WITH_TZ,
now,
)
from sientia_do.utils.formatters import create_sample_dict
from laborious.utils.models.minio_dataframe_payload import MinioDataFramePayload
from laborious.utils.dataframe_debug import build_dataframe_debug_message
from laborious.utils.repository.minio_manager import MinioManager
from laborious.utils.repository.model_repository import MLFlowRepository
class MLFlow(MinioManager):
"""
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
"""
_MAX_DEBUG_DATAFRAME_ROWS = 100
def __init__(
self,
mlflow_host: str,
mlflow_port: int,
mlflow_username: str,
mlflow_password: str,
minio_repository: MinioRepository | None = None,
logger: Logger | None = None,
notification_handler: NotificationHandler | None = None,
metrics_controller: MetricsController | None = None,
):
"""
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
"""
MinioManager.__init__(
self, minio_repository, logger, notification_handler, metrics_controller
)
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,
notification_handler,
metrics_controller,
)
def close(self) -> None:
"""
Close the MLFlow activity and clean up resources.
"""
MinioManager.close(self)
def __del__(self):
self.close()
def _debug_dataframe(self, message: str, data: Any, metadata: dict[str, Any]) -> None:
"""
Log dataframe content only when row count is below the configured threshold
Args:
- message (str): Base log message to identify the dataframe in logs
- data (Any): Dataframe-like object expected to expose shape and to_csv
- metadata (dict[str, Any]): Workflow metadata for contextual logging
"""
self.debug(
build_dataframe_debug_message(
message=message,
data=data,
max_rows=self._MAX_DEBUG_DATAFRAME_ROWS,
),
metadata,
)
@activity.defn(name='request_transform')
async def request_transform(self, input_data: dict[str, Any]) -> MinioDataFramePayload:
"""
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)
payload = MinioDataFramePayload.from_dict(input_data['data'])
data = await payload.retrieve(self.minio_repository, metadata)
model_name = input_data['model_name']
model_config = input_data.get('model_config', {})
self._debug_dataframe('Raw input data:', data, 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.columns.name = None
data.index.name = None
data['timestamp'] = data.index
self._debug_dataframe('Processed input data:', data, metadata)
# Request transformation from MLFlow model
response_data = await 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)
if not response_data.get('success', False):
return await MinioDataFramePayload.from_dataframe(
dataframe=None,
minio_repo=self.minio_repository,
model_name=model_name,
operation='transform',
status=response_data,
workflow_metadata=metadata,
last_timestamp=payload.last_timestamp,
logger=self.logger,
)
return await MinioDataFramePayload.from_dataframe(
dataframe=response_data['content'],
minio_repo=self.minio_repository,
model_name=model_name,
operation='transform',
workflow_metadata=metadata,
status={
'success': True,
},
last_timestamp=payload.last_timestamp,
logger=self.logger,
)
@activity.defn(name='request_predict')
async def request_predict(self, input_data: dict[str, Any]) -> MinioDataFramePayload:
"""
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)
payload = MinioDataFramePayload.from_dict(input_data['data'])
data = await payload.retrieve(self.minio_repository, metadata)
model_name = input_data['model_name']
model_config = input_data.get('model_config', {})
self._debug_dataframe('Input data for prediction:', data, 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 = await 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)
if not response_data.get('success', False):
return await MinioDataFramePayload.from_dataframe(
dataframe=None,
minio_repo=self.minio_repository,
model_name=model_name,
operation='predict',
status=response_data,
workflow_metadata=metadata,
last_timestamp=payload.last_timestamp,
logger=self.logger,
)
return await MinioDataFramePayload.from_dataframe(
dataframe=response_data['content'],
minio_repo=self.minio_repository,
model_name=model_name,
operation='predict',
workflow_metadata=metadata,
status={
'success': True,
},
last_timestamp=payload.last_timestamp,
logger=self.logger,
)
@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
"""
if self.minio_repository is None:
raise ValueError('Minio repository not initialized')
metadata = input_data['metadata']
try:
# Payload-based retrain input (inline dict or MinIO offloaded).
payload = MinioDataFramePayload.from_dict(input_data['data'])
data = await payload.retrieve(self.minio_repository, metadata)
except Exception as e:
trace = traceback.format_exc()
await self.send_notification_async(
metadata=metadata,
notification_id='ERROR_LOADING_RETRAIN_DATA',
message=f'Error loading retrain data: {e}',
block='retrain_model',
level=NotificationLevel.ERROR,
attachment_content=trace,
)
self.error(trace, metadata)
return {
'success': False,
'message': f'Error loading retrain data: {e}',
'traceback': trace,
'timestamp': now().strftime(DATETIME_FORMAT_MS_WITH_TZ),
}
self.debug(f'Retrain data loaded successfully: shape {data.shape}', metadata)
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)
# Sort by created_at in descending order and keep first occurrence of each variable/timestamp pair
if 'created_at' in data.columns:
data = data.sort_values('created_at', ascending=False).drop_duplicates(
subset=['variable', 'timestamp'], keep='first'
)
else:
data = data.drop_duplicates(subset=['variable', 'timestamp'], keep='first')
data.drop(columns=['model_id'], inplace=True, errors='ignore')
data.drop(columns=['created_at'], inplace=True, errors='ignore')
# 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
data['timestamp'] = data.index
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.columns.name = None
retrain_output = await 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']
await self.send_notification_async(
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 = await self.model_monitoring_repository.update_production_model(
experiment=experiment, model_name=model_name, metadata=metadata
)
self.info(f'Production model {model_name} updated successfully', metadata)
return response
except Exception as e:
trace = traceback.format_exc()
await self.send_notification_async(
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='get_reference_data')
async def get_reference_data(self, input_data: dict[str, Any]) -> list[dict] | None:
"""
Get reference data from the MLflow Model Registry.
This method retrieves evaluation reference data stored as artifacts in the
MLflow Model Registry. The reference data is typically used for model
drift detection, performance comparison, and quality validation. The method
loads the data from a CSV artifact file and formats timestamps for
consistent processing.
The method handles:
1. Loading evaluation data artifact from MLflow Model Registry
2. Timestamp parsing and formatting for consistency
3. Data conversion to dictionary format for workflow consumption
4. Graceful handling of missing reference data
Args:
input_data (dict): Input data containing:
- metadata (dict): Workflow execution metadata
- model_name (str): Name of the MLFlow model to get reference data from
Returns:
list[dict[Hashable, Any]] | None: Reference data from the MLflow Model Registry
as a list of dictionaries. Returns None if reference data is not found
or if the artifact does not exist.
Raises:
Exception: If artifact loading fails or encounters errors during processing
"""
metadata = input_data['metadata']
model_name = input_data['model_name']
artifact = 'evaluation_data.csv'
reference_data = await self.model_monitoring_repository.load_artifact_dataframe(
model_name=model_name, artifact_path=artifact, metadata=metadata
)
if reference_data is None:
self.warning(f'Reference data not found for model {model_name}', metadata)
return None
reference_data['timestamp'] = to_datetime(reference_data['timestamp'])
reference_data['timestamp'] = reference_data['timestamp'].dt.strftime(DATETIME_FORMAT)
return reference_data.to_dict(orient='records')