""" Training activities for ML model training operations. This module provides activities for training machine learning models. The activity extends BaseActivity and receives pre-downloaded files and raises `ModelTrainingError` when training fails. """ from sientia_model.wrappers.sientia_model import SientiaModel from temporalio import activity, workflow with workflow.unsafe.imports_passed_through(): import traceback from typing import Any import mlflow 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.observability.sientia_monitoring import SientiaMonitoring from sientia_do.repository.minio_repository_sync import MinioRepository from sientia_model.model_repository.mlflow_repository import SientiaMLflowRepository from sientia_model.model_repository.plugin_store import PluginStore from model_manager.utils.models.train_model_params import TrainModelParams from model_manager.utils.models.train_model_result import TrainModelResult from model_manager.utils.repository.data_manager_repository import DataManagerRepository class Training(SientiaMonitoring): """ Activity for ML model training operations. This activity extends SientiaMonitoring and handles machine learning model training with comprehensive error handling. It receives pre-downloaded files from the workflow and raises `ModelTrainingError` on failure so the workflow can map the correct experiment status. """ def __init__( self, mlflow_repository: SientiaMLflowRepository, plugin_store: PluginStore, minio_repository: MinioRepository, logger: Logger, notification_handler: NotificationHandler, metrics_controller: MetricsController, ): """ Initialize Training activity. Args: logger: Logger instance for observability notification_handler: Handler for sending notifications """ SientiaMonitoring.__init__(self, logger, notification_handler, metrics_controller) self.data_manager_repository = DataManagerRepository(logger) self.mlflow_repository = mlflow_repository self.plugin_store = plugin_store self.minio_repository = minio_repository @activity.defn(name='load_model_metadata') def load_model_metadata(self, input_data: dict[str, Any]) -> dict[str, Any]: """ Load model metadata/schemas from the model store. This activity is responsible for fetching model metadata/schemas from the model store index and extracting a serializable `model_metadata` dict that `TrainModelParams.validate_business_rules()` depends on. Args: input_data: Workflow input at the same level as `validate_train_params`, including at least `model_name` and the fields required by `TrainModelParams.from_dict` to build wrapper kwargs. Return: dict[str, Any]: Updated `input_data` containing `input_data['model_metadata']`. """ metadata = input_data.get('metadata', {}) self.info(f'Loading model metadata for {input_data}', metadata) try: train_params = TrainModelParams.from_dict(input_data) model_metadata = self.plugin_store.get_model_index( model_type=train_params.model_type, metadata=metadata, ) train_params.model_metadata = model_metadata self.info(f'Model metadata loaded successfully for {input_data}', metadata) self.debug(f'Model metadata: {model_metadata}', metadata) return train_params.to_dict() except Exception as exc: trace = traceback.format_exc() self.send_notification( metadata=metadata, notification_id='LOAD_MODEL_METADATA_ERROR', message=f'Error loading model metadata: {str(exc)}', block='load_model_metadata', level=NotificationLevel.ERROR, attachment_content=trace, ) raise @activity.defn(name='validate_train_params') def validate_train_params(self, input_data: dict[str, Any]) -> dict[str, Any]: """ Validate and convert training parameters from dict to TrainModelParams. This activity validates the input training parameters and converts them to a TrainModelParams object. Args: input_data: Training parameters and metadata at the same level Required keys: - metadata (dict): Workflow execution metadata - All TrainModelParams fields (experiment_run_id, target_variable, etc.) Returns: dict[str, Any]: Validated and converted training parameters as dictionary Raises: Exception: If validation fails (after sending notification) """ metadata = input_data.get('metadata', {}) self.info(f'Validating training parameters for {input_data}', metadata) try: train_params = TrainModelParams.from_dict(input_data) train_params.validate_business_rules() self.info( f'Training parameters validated successfully - ' f'Target: {train_params.target_variable}, ' f'Experiment: {train_params.experiment_name}', metadata, ) self.debug( f'Training parameters validated successfully: {train_params.to_dict()}', metadata ) return train_params.to_dict() except Exception as e: error_msg = f'Error validating training parameters: {str(e)}' trace = traceback.format_exc() self.send_notification( metadata=metadata, notification_id='VALIDATE_TRAIN_PARAMS_ERROR', message=error_msg, block='validate_train_params', level=NotificationLevel.ERROR, attachment_content=trace, ) raise @activity.defn(name='train_model') def train_model(self, input_data: dict[str, Any]) -> dict[str, Any]: """ Train a machine learning model. This activity orchestrates the ML training pipeline: 1. Validate input parameters. 2. Prepare data via DataManagerRepository. 3. Train the model and compute metrics. Args: input_data: Training configuration containing: - metadata (dict): Workflow execution metadata. - uploaded_file (BytesIO): Training data already downloaded from MinIO. - train_params (dict): Training parameters. Returns: dict[str, Any]: Serializable summary (run identifiers, run_dir for cleanup, regression metrics). Raises: ValueError: If input validation fails. Exception: If training fails (after sending notification). """ metadata = input_data.get('metadata') train_params = TrainModelParams.from_dict(input_data['train_params']) self.info('Starting train_model process', metadata) try: # Download training file bytes from MinIO self.info( f'Downloading training file from MinIO for {train_params.file_name}', metadata ) train_bytes = self.minio_repository.download_file( object_name=train_params.file_name, bucket=train_params.bucket_name, metadata=metadata, ) # Download optional validation file bytes from the same bucket val_bytes: bytes | None = None validation_name = train_params.val_file_name if validation_name is not None: self.info(f'Downloading validation file from MinIO for {validation_name}', metadata) val_bytes = self.minio_repository.download_file( object_name=validation_name, bucket=train_params.bucket_name, metadata=metadata, ) self.info(f'Preparing training data for {train_params.file_name}', metadata) train_result = self.data_manager_repository.prepare_training_data( train_file_bytes=train_bytes, validation_file_bytes=val_bytes, params=train_params, metadata=metadata, ) self.info(f'Getting model wrapper for {train_params.model_type}', metadata) wrapper = self.plugin_store.get_model( model_type=train_params.model_type, force_download=False, opt_params=train_params.opt_params or {}, model_kwargs=train_params.model_kwargs or {}, data_model_kwargs=train_params.data_model_kwargs or {}, metadata=metadata, ) if self.logger is not None: wrapper.logger = self.logger.base_logger self.info(f'Training model for {train_params.model_type}', metadata) train_data = train_result.train_data val_data = train_result.val_data self.debug( f'train_model prepared data (head 10):\ntrain:\n{train_data.head(10).to_string()}' f'\nval:\n{val_data.head(10).to_string()}', metadata, ) wrapper.train( train_data=train_data, val_data=val_data, target=train_params.target_variable, ) self.info( f'Generating predictions using the trained wrapper for {train_params.model_type}', metadata, ) # Generate predictions using the trained wrapper transformed_train, _ = wrapper.transform(train_data) transformed_val, _ = wrapper.transform(val_data) self.debug( f'train_model transform (head 10):\ntrain:\n{transformed_train.head(10).to_string()}' f'\nval:\n{transformed_val.head(10).to_string()}', metadata, ) y_train_pred_df, _ = wrapper.predict({}, transformed_train) y_val_pred_df, _ = wrapper.predict({}, transformed_val) self.debug( f'train_model predict (head 10):\ntrain:\n{y_train_pred_df.head(10).to_string()}' f'\nval:\n{y_val_pred_df.head(10).to_string()}', metadata, ) y_train_pred_df.sort_index(inplace=True, ascending=False) y_val_pred_df.sort_index(inplace=True, ascending=False) train_result.y_train_pred = y_train_pred_df train_result.y_pred = y_val_pred_df self.info(f'Computing regression metrics for {train_params.model_type}', metadata) train_result = self.data_manager_repository.compute_regression_metrics( train_result, wrapper, metadata=metadata, ) self.info(f'Starting MLflow run for {train_params.model_type}', metadata) with self.mlflow_repository.start_run( model_name=train_params.model_name, run_name=train_result.run_name, experiment_name=train_result.experiment_name, tags=None, metadata=metadata, ) as run_info: train_result.run_id = run_info.run_id self._persist_training_artifacts(train_result, train_params, wrapper, metadata) return { 'run_name': train_result.run_name, 'experiment_name': train_result.experiment_name, 'run_id': train_result.run_id, 'run_dir': train_result.run_dir, } except Exception as e: # noqa: BLE001 error_msg = f'Error training model - error: {str(e)}' trace = traceback.format_exc() self.send_notification( metadata=metadata or {}, notification_id='TRAIN_MODEL_ERROR', message=error_msg, block='train_model', level=NotificationLevel.ERROR, attachment_content=trace, ) raise e def _persist_training_artifacts( self, train_result: TrainModelResult, train_params: TrainModelParams, wrapper: SientiaModel, metadata: dict[str, Any] | None, ) -> None: self.info(f'Generating report for {train_params.model_type}', metadata) train_result = self.data_manager_repository.generate_report( train_result, metadata=metadata, ) if ( train_result.report_path is None or train_result.train_data_path is None or train_result.test_data_path is None ): raise ValueError('Report path, train data path, or test data path is not set') self.info(f'Storing model for {train_params.model_type}', metadata) wrapper._input_example = None wrapper.store_model(name=train_params.model_name) self._log_regression_metrics_as_params(train_result) self.info(f'Logging artifacts for {train_params.model_type}', metadata) mlflow.log_artifact(train_result.report_path) mlflow.log_artifact(train_result.train_data_path) mlflow.log_artifact(train_result.test_data_path) if train_result.equation_path is not None: mlflow.log_artifact(train_result.equation_path) def _log_regression_metrics_as_params(self, train_result: TrainModelResult) -> None: """ Persist computed regression metrics as MLflow params. Args: train_result: Training output containing computed regression metrics. """ metric_params = { 'mse_val': train_result.mse_val, 'mae_val': train_result.mae_val, 'r2_val': train_result.r2_val, } for key, value in metric_params.items(): if value is not None: mlflow.log_param(key, value) @activity.defn(name='cleanup_resources') def cleanup_resources(self, input_data: dict[str, Any]) -> None: """ Cleanup temporary resources created during training. Args: input_data: Cleanup configuration containing: - metadata (dict): Workflow execution metadata. - run_dir (str): Temporary directory to remove. Raises: Exception: If cleanup fails (after sending notification). """ metadata = input_data.get('metadata', {}) run_dir = input_data.get('run_dir', '') try: self.data_manager_repository.cleanup_run_directory(run_dir, metadata) except Exception as e: # noqa: BLE001 error_msg = f'Error cleaning up resources - Run directory: {run_dir}, Error: {str(e)}' trace = traceback.format_exc() self.send_notification( metadata=metadata, notification_id='CLEANUP_RESOURCES_ERROR', message=error_msg, block='cleanup_resources', level=NotificationLevel.ERROR, attachment_content=trace, ) raise