Merge pull request #5 from Aignosi/feature/SIENTIAPDE-1251
Implement ML Model Training Activity and Repository (SIENTIAPDE-1251)
This commit is contained in:
@@ -154,6 +154,12 @@ The Model Manager system uses a Temporal-based workflow architecture with clear
|
|||||||
- Support for three update types: STATUS, STATUS_WITH_ERROR, MODEL_SAVED
|
- Support for three update types: STATUS, STATUS_WITH_ERROR, MODEL_SAVED
|
||||||
- Automatic error message truncation (1024 chars)
|
- Automatic error message truncation (1024 chars)
|
||||||
- Connection pooling and retry logic via Postgres base class
|
- Connection pooling and retry logic via Postgres base class
|
||||||
|
- **Training**: ML model training operations (standalone activity, composition pattern)
|
||||||
|
- Unified `train_model()` method for complete training pipeline
|
||||||
|
- Receives pre-downloaded files (BytesIO) to avoid memory leaks
|
||||||
|
- Returns success/failure status with TrainModelResult or error message
|
||||||
|
- No exception raising on failure - allows workflow to handle errors gracefully
|
||||||
|
- Integration with TrainingRepository for business logic separation
|
||||||
- **Gates**: Data quality validation and filtering mechanisms
|
- **Gates**: Data quality validation and filtering mechanisms
|
||||||
- **MLFlow**: Model transformation and prediction operations
|
- **MLFlow**: Model transformation and prediction operations
|
||||||
- **MinIO**: Object storage operations for file management
|
- **MinIO**: Object storage operations for file management
|
||||||
|
|||||||
@@ -10,9 +10,10 @@ with workflow.unsafe.imports_passed_through():
|
|||||||
from model_manager.activities.gates import Gates
|
from model_manager.activities.gates import Gates
|
||||||
from model_manager.activities.minio import MinIO
|
from model_manager.activities.minio import MinIO
|
||||||
from model_manager.activities.mlflow import MLFlow
|
from model_manager.activities.mlflow import MLFlow
|
||||||
|
from model_manager.activities.training import Training
|
||||||
|
|
||||||
|
|
||||||
class Activities(ExperimentTracking, MLFlow, MinIO, Gates):
|
class Activities(ExperimentTracking, MLFlow, MinIO, Gates, Training):
|
||||||
"""
|
"""
|
||||||
Main activities orchestrator for the Model Manager system.
|
Main activities orchestrator for the Model Manager system.
|
||||||
|
|
||||||
@@ -25,6 +26,7 @@ class Activities(ExperimentTracking, MLFlow, MinIO, Gates):
|
|||||||
- MLFlow: Model inference and transformation operations
|
- MLFlow: Model inference and transformation operations
|
||||||
- MinIO: Object storage operations (file upload/download/delete)
|
- MinIO: Object storage operations (file upload/download/delete)
|
||||||
- Gates: Data quality validation and filtering mechanisms
|
- Gates: Data quality validation and filtering mechanisms
|
||||||
|
- Training: ML model training operations (extends BaseActivity)
|
||||||
|
|
||||||
Attributes:
|
Attributes:
|
||||||
postgres_config (dict): PostgreSQL connection configuration
|
postgres_config (dict): PostgreSQL connection configuration
|
||||||
@@ -102,6 +104,8 @@ class Activities(ExperimentTracking, MLFlow, MinIO, Gates):
|
|||||||
|
|
||||||
Gates.__init__(self, logger=logger, notification_handler=notification_handler)
|
Gates.__init__(self, logger=logger, notification_handler=notification_handler)
|
||||||
|
|
||||||
|
Training.__init__(self, logger=logger, notification_handler=notification_handler)
|
||||||
|
|
||||||
async def shutdown(self):
|
async def shutdown(self):
|
||||||
"""
|
"""
|
||||||
Gracefully shutdown all activities and clean up resources.
|
Gracefully shutdown all activities and clean up resources.
|
||||||
|
|||||||
166
model_manager/activities/training.py
Normal file
166
model_manager/activities/training.py
Normal file
@@ -0,0 +1,166 @@
|
|||||||
|
"""
|
||||||
|
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
|
||||||
|
to return success/failure status without raising exceptions.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from temporalio import activity, workflow
|
||||||
|
|
||||||
|
with workflow.unsafe.imports_passed_through():
|
||||||
|
import traceback
|
||||||
|
from io import BytesIO
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
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 model_manager.utils.models.train_model_params import TrainModelParams
|
||||||
|
from model_manager.utils.repository.training_repository import TrainingRepository
|
||||||
|
|
||||||
|
|
||||||
|
class Training(BaseActivity):
|
||||||
|
"""
|
||||||
|
Activity for ML model training operations.
|
||||||
|
|
||||||
|
This activity extends BaseActivity and handles machine learning model
|
||||||
|
training with comprehensive error handling. It receives pre-downloaded
|
||||||
|
files from the workflow and returns success/failure status without
|
||||||
|
raising exceptions.
|
||||||
|
|
||||||
|
Attributes:
|
||||||
|
logger (Logger): Logger instance for observability (inherited from BaseActivity)
|
||||||
|
notification_handler (NotificationHandler): Handler for sending notifications (inherited)
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
logger: Logger,
|
||||||
|
notification_handler: NotificationHandler,
|
||||||
|
):
|
||||||
|
"""
|
||||||
|
Initialize Training activity.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
logger: Logger instance for observability
|
||||||
|
notification_handler: Handler for sending notifications
|
||||||
|
"""
|
||||||
|
BaseActivity.__init__(self, logger, notification_handler, set_error_counter=True)
|
||||||
|
self.training_repository = TrainingRepository(logger)
|
||||||
|
|
||||||
|
@activity.defn(name='train_model')
|
||||||
|
async def train_model(self, input_data: dict[str, Any]) -> dict[str, Any]:
|
||||||
|
"""
|
||||||
|
Train a machine learning model with comprehensive error handling.
|
||||||
|
|
||||||
|
This activity orchestrates the complete ML training pipeline:
|
||||||
|
1. Validates input parameters
|
||||||
|
2. Trains the model using TrainingRepository
|
||||||
|
3. Performs post-training calculations
|
||||||
|
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 training operation
|
||||||
|
Required keys:
|
||||||
|
- metadata (dict): Workflow execution metadata
|
||||||
|
- uploaded_file (BytesIO): Training data file (already downloaded from MinIO)
|
||||||
|
- train_params (dict): Training parameters (converted to TrainModelParams)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
dict: Training result with the following structure:
|
||||||
|
{
|
||||||
|
'success': bool, # True if training succeeded, False otherwise
|
||||||
|
'result': TrainModelResult | None, # Training result if success=True
|
||||||
|
'error_message': str | None # Error message if success=False
|
||||||
|
}
|
||||||
|
|
||||||
|
Example:
|
||||||
|
# Successful training
|
||||||
|
result = await train_model({
|
||||||
|
'metadata': {'workflow_id': 'train-123', 'experiment_run_id': 456},
|
||||||
|
'uploaded_file': BytesIO(csv_data),
|
||||||
|
'train_params': {
|
||||||
|
'experiment_run_id': 456,
|
||||||
|
'target_variable': 'price',
|
||||||
|
'variable_columns': ['feature1', 'feature2'],
|
||||||
|
'train_size': 80,
|
||||||
|
'shuffle': True,
|
||||||
|
'use_scaler': True,
|
||||||
|
# ... other TrainModelParams fields
|
||||||
|
}
|
||||||
|
})
|
||||||
|
# Returns: {'success': True, 'result': TrainModelResult(...), 'error_message': None}
|
||||||
|
|
||||||
|
# Failed training
|
||||||
|
# Returns: {'success': False, 'result': None, 'error_message': 'Error details...'}
|
||||||
|
"""
|
||||||
|
metadata = input_data.get('metadata', {})
|
||||||
|
uploaded_file = input_data['uploaded_file']
|
||||||
|
train_params_dict = input_data['train_params']
|
||||||
|
|
||||||
|
try:
|
||||||
|
self.info(
|
||||||
|
f'Starting model training for target: {train_params_dict.get("target_variable")}',
|
||||||
|
metadata,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Convert dict to TrainModelParams
|
||||||
|
train_params = TrainModelParams.from_dict(train_params_dict)
|
||||||
|
|
||||||
|
# Validate uploaded_file is BytesIO
|
||||||
|
if not isinstance(uploaded_file, BytesIO):
|
||||||
|
raise ValueError(f'uploaded_file must be BytesIO, got {type(uploaded_file)}')
|
||||||
|
|
||||||
|
# Step 1: Train the model
|
||||||
|
self.info('Training model with TrainingRepository', metadata)
|
||||||
|
train_result = self.training_repository.train(uploaded_file, train_params)
|
||||||
|
|
||||||
|
# Step 2: Perform post-training calculations
|
||||||
|
self.info('Performing post-training calculations', metadata)
|
||||||
|
final_result = self.training_repository.after_train_calculation(
|
||||||
|
train_params, train_result
|
||||||
|
)
|
||||||
|
|
||||||
|
self.info(
|
||||||
|
f'Model training completed successfully - '
|
||||||
|
f'MSE: {final_result.mse_val}, MAE: {final_result.mae_val}, R²: {final_result.r2_val}',
|
||||||
|
metadata,
|
||||||
|
)
|
||||||
|
|
||||||
|
return {
|
||||||
|
'success': True,
|
||||||
|
'result': final_result,
|
||||||
|
'error_message': None,
|
||||||
|
}
|
||||||
|
|
||||||
|
except Exception as e: # noqa: BLE001
|
||||||
|
error_msg = f'Error training model - Target: {train_params_dict.get("target_variable", "unknown")}, Error: {str(e)}'
|
||||||
|
trace = traceback.format_exc()
|
||||||
|
|
||||||
|
# Send notification (MongoDB)
|
||||||
|
self.send_notification(
|
||||||
|
metadata=metadata,
|
||||||
|
notification_id='TRAIN_MODEL_ERROR',
|
||||||
|
message=error_msg,
|
||||||
|
block='train_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),
|
||||||
|
}
|
||||||
242
model_manager/utils/repository/training_repository.py
Normal file
242
model_manager/utils/repository/training_repository.py
Normal file
@@ -0,0 +1,242 @@
|
|||||||
|
"""
|
||||||
|
Training repository for ML model training operations.
|
||||||
|
|
||||||
|
This module provides the core training logic for machine learning models,
|
||||||
|
including data preprocessing, model training, and post-training calculations.
|
||||||
|
Migrated from laborious/utils/train_model_utils.py.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from io import BytesIO
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
import pandas as pd
|
||||||
|
from sientia.linear_models import LinearRegressionModel
|
||||||
|
from sientia.metrics import mae, mse, r2
|
||||||
|
from sientia.preprocessing import DataPreprocessor
|
||||||
|
from sientia.utils import split_train_test
|
||||||
|
from sientia_do.observability.logger import Logger
|
||||||
|
from sientia_do.operations.df_preprocessor import load_data
|
||||||
|
from sientia_do.operations.normalization import MinMaxScaler, Z_Scaler
|
||||||
|
|
||||||
|
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||||
|
from model_manager.utils.models.train_model_result import TrainModelResult
|
||||||
|
|
||||||
|
|
||||||
|
class TrainingRepository:
|
||||||
|
"""
|
||||||
|
Repository for machine learning model training operations.
|
||||||
|
|
||||||
|
This class encapsulates the core logic for training ML models, migrated from
|
||||||
|
laborious/utils/train_model_utils.py. Follows the same pattern as MLFlowRepository
|
||||||
|
with instance methods and logger integration.
|
||||||
|
|
||||||
|
Attributes:
|
||||||
|
logger (Logger): Logger instance for observability and debugging
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, logger: Logger):
|
||||||
|
"""
|
||||||
|
Initialize TrainingRepository with logger.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
logger: Logger instance for observability
|
||||||
|
"""
|
||||||
|
self.logger = logger
|
||||||
|
|
||||||
|
def train(self, uploaded_file: BytesIO, params: TrainModelParams) -> TrainModelResult:
|
||||||
|
"""
|
||||||
|
Train a machine learning model using the provided file and parameters.
|
||||||
|
|
||||||
|
This method orchestrates the training pipeline:
|
||||||
|
1. Load data from BytesIO file
|
||||||
|
2. Initialize and fit data preprocessor
|
||||||
|
3. Transform data and validate
|
||||||
|
4. Split into train/test sets
|
||||||
|
5. Initialize scaler dictionary
|
||||||
|
6. Train LinearRegression model
|
||||||
|
|
||||||
|
Args:
|
||||||
|
uploaded_file: BytesIO object containing training data (CSV format)
|
||||||
|
params: Training parameters (TrainModelParams)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
TrainModelResult: Object containing trained model, processed data,
|
||||||
|
train/test splits, and scaler dictionary
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
ValueError: If transformed data is empty
|
||||||
|
Exception: If data loading, preprocessing, or training fails
|
||||||
|
"""
|
||||||
|
# Load data from BytesIO
|
||||||
|
data = load_data(uploaded_file, params.line_separator, params.decimal_separator)
|
||||||
|
|
||||||
|
# Initialize and fit data preprocessor
|
||||||
|
process_data = self.init_data_preprocessor(params)
|
||||||
|
process_data.fit(data)
|
||||||
|
data_view = process_data.transform(data)
|
||||||
|
|
||||||
|
# Validate transformed data
|
||||||
|
if len(data_view) <= 0:
|
||||||
|
raise ValueError('Data view is empty after transformation')
|
||||||
|
|
||||||
|
# Split data into train/test sets
|
||||||
|
x_train, x_test, y_train, y_test = split_train_test(
|
||||||
|
data_view[params.variable_columns],
|
||||||
|
data_view[params.target_variable],
|
||||||
|
train_size=params.train_size / 100,
|
||||||
|
shuffle=params.shuffle,
|
||||||
|
random_state=42,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Prepare training data
|
||||||
|
data_train = pd.concat([x_train, y_train], axis=1)
|
||||||
|
scaler_dict = self.init_scaler_dict(process_data, params)
|
||||||
|
|
||||||
|
# Create and train linear regression model
|
||||||
|
regr = LinearRegressionModel(
|
||||||
|
target_variable=params.target_variable,
|
||||||
|
variable_columns=params.variable_columns,
|
||||||
|
)
|
||||||
|
regr.fit(data_train)
|
||||||
|
|
||||||
|
# Return training result
|
||||||
|
return TrainModelResult(
|
||||||
|
params=params,
|
||||||
|
process_data=process_data,
|
||||||
|
x_train=x_train,
|
||||||
|
x_test=x_test,
|
||||||
|
y_train=y_train,
|
||||||
|
y_test=y_test,
|
||||||
|
regr=regr,
|
||||||
|
scaler_dict=scaler_dict,
|
||||||
|
)
|
||||||
|
|
||||||
|
def init_scaler_dict(self, process_data: DataPreprocessor, params: TrainModelParams) -> dict:
|
||||||
|
"""
|
||||||
|
Initialize dictionary containing scaling parameters for features and target.
|
||||||
|
|
||||||
|
This method extracts scaling parameters from the fitted scaler to enable
|
||||||
|
denormalization of predictions and debugging of the normalization process.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
process_data: Fitted DataPreprocessor object with scaler
|
||||||
|
params: Training parameters including scaler configuration
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
dict: Scaling parameters for each feature and target variable.
|
||||||
|
Structure depends on scaler type:
|
||||||
|
- MinMaxScaler: {'feature': {'min': float, 'max': float}, ...}
|
||||||
|
- Z_Scaler: Dictionary from scaler.create_dict()
|
||||||
|
- Empty dict: If no scaler is used
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
AttributeError: If scaler doesn't have expected attributes
|
||||||
|
"""
|
||||||
|
scaler_dict = {}
|
||||||
|
|
||||||
|
if params.use_scaler:
|
||||||
|
scaler = process_data.get_scaler()
|
||||||
|
|
||||||
|
if isinstance(scaler, MinMaxScaler):
|
||||||
|
# Extract min/max for each feature
|
||||||
|
for i, col in enumerate(params.variable_columns):
|
||||||
|
scaler_dict[col] = {'min': scaler.x_min[i], 'max': scaler.x_max[i]}
|
||||||
|
|
||||||
|
# Extract min/max for target variable
|
||||||
|
scaler_dict[params.target_variable] = {
|
||||||
|
'min': scaler.y_min,
|
||||||
|
'max': scaler.y_max,
|
||||||
|
}
|
||||||
|
|
||||||
|
elif isinstance(scaler, Z_Scaler):
|
||||||
|
scaler_dict = scaler.create_dict()
|
||||||
|
|
||||||
|
return scaler_dict
|
||||||
|
|
||||||
|
def after_train_calculation(
|
||||||
|
self, params: TrainModelParams, tmr: TrainModelResult
|
||||||
|
) -> TrainModelResult:
|
||||||
|
"""
|
||||||
|
Perform post-training calculations: predictions, denormalization, and metrics.
|
||||||
|
|
||||||
|
This method completes the training pipeline by:
|
||||||
|
1. Making predictions on test set
|
||||||
|
2. Denormalizing all data (if scaler was used)
|
||||||
|
3. Reordering data by index
|
||||||
|
4. Calculating evaluation metrics (MSE, MAE, R²)
|
||||||
|
|
||||||
|
Args:
|
||||||
|
params: Training parameters used during model training
|
||||||
|
tmr: Result object from training
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
TrainModelResult: Updated result with predictions, denormalized data,
|
||||||
|
and metrics (mse_val, mae_val, r2_val)
|
||||||
|
"""
|
||||||
|
# Make predictions on test set
|
||||||
|
tmr.y_pred = tmr.regr.predict(tmr.x_test)
|
||||||
|
|
||||||
|
# Denormalize data if scaler was used
|
||||||
|
if params.use_scaler:
|
||||||
|
scaler = tmr.process_data.get_scaler()
|
||||||
|
|
||||||
|
# Denormalize features
|
||||||
|
for col in params.variable_columns:
|
||||||
|
tmr.x_train[col] = scaler.denormalize_single_input(tmr.x_train[col], col)
|
||||||
|
tmr.x_test[col] = scaler.denormalize_single_input(tmr.x_test[col], col)
|
||||||
|
|
||||||
|
# Denormalize target variable
|
||||||
|
tmr.y_train = scaler.denormalize_single_input(tmr.y_train, params.target_variable)
|
||||||
|
tmr.y_test = scaler.denormalize_single_input(tmr.y_test, params.target_variable)
|
||||||
|
tmr.y_pred = scaler.denormalize_predictions(tmr.y_pred, params.target_variable)
|
||||||
|
|
||||||
|
# Add index to predictions
|
||||||
|
tmr.y_pred = pd.Series(tmr.y_pred, index=tmr.y_test.index)
|
||||||
|
tmr.y_pred.name = f'{params.target_variable}_pred'
|
||||||
|
|
||||||
|
# Reorder all data by index
|
||||||
|
tmr.x_train = tmr.x_train.sort_index()
|
||||||
|
tmr.x_test = tmr.x_test.sort_index()
|
||||||
|
tmr.y_train = tmr.y_train.sort_index()
|
||||||
|
tmr.y_test = tmr.y_test.sort_index()
|
||||||
|
tmr.y_pred = tmr.y_pred.sort_index()
|
||||||
|
|
||||||
|
# Calculate evaluation metrics
|
||||||
|
tmr.mse_val = round(
|
||||||
|
mse(tmr.y_test.astype(np.float64), tmr.y_pred.astype(np.float64)),
|
||||||
|
2,
|
||||||
|
)
|
||||||
|
tmr.mae_val = round(
|
||||||
|
mae(tmr.y_test.astype(np.float64), tmr.y_pred.astype(np.float64)),
|
||||||
|
2,
|
||||||
|
)
|
||||||
|
tmr.r2_val = round(r2(tmr.y_test.astype(np.float64), tmr.y_pred.astype(np.float64)), 2)
|
||||||
|
|
||||||
|
return tmr
|
||||||
|
|
||||||
|
def init_data_preprocessor(self, params: TrainModelParams) -> DataPreprocessor:
|
||||||
|
"""
|
||||||
|
Initialize DataPreprocessor with training parameters.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
params: Training parameters containing preprocessor configuration
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
DataPreprocessor: Configured preprocessor ready for fitting
|
||||||
|
"""
|
||||||
|
# Create lag dictionaries for each variable
|
||||||
|
lag_train_dict = dict.fromkeys(params.variable_columns, params.lag_train)
|
||||||
|
lag_val_dict = dict.fromkeys(params.variable_columns, params.lag_val)
|
||||||
|
|
||||||
|
return DataPreprocessor(
|
||||||
|
target_variable=params.target_variable,
|
||||||
|
input_columns=params.variable_columns,
|
||||||
|
lag_train=lag_train_dict,
|
||||||
|
lag_transform=lag_val_dict,
|
||||||
|
static_threshold=1 if params.rem_static_win else None,
|
||||||
|
low_lim=params.low_lim,
|
||||||
|
upp_lim=params.upp_lim,
|
||||||
|
window=params.window,
|
||||||
|
scaler_name='Standard Scaler' if params.use_scaler else 'None',
|
||||||
|
ar_var=params.target_variable if params.include_ar else None,
|
||||||
|
)
|
||||||
@@ -6,14 +6,20 @@ from model_manager.activities.activities import Activities
|
|||||||
from model_manager.activities.experiment_tracking import ExperimentTracking
|
from model_manager.activities.experiment_tracking import ExperimentTracking
|
||||||
from model_manager.activities.gates import Gates
|
from model_manager.activities.gates import Gates
|
||||||
from model_manager.activities.mlflow import MLFlow
|
from model_manager.activities.mlflow import MLFlow
|
||||||
|
from model_manager.activities.training import Training
|
||||||
|
|
||||||
|
|
||||||
@patch('model_manager.activities.activities.ExperimentTracking.__init__')
|
@patch('model_manager.activities.activities.ExperimentTracking.__init__')
|
||||||
@patch('model_manager.activities.activities.MLFlow.__init__')
|
@patch('model_manager.activities.activities.MLFlow.__init__')
|
||||||
@patch('model_manager.activities.activities.MinIO.__init__')
|
@patch('model_manager.activities.activities.MinIO.__init__')
|
||||||
@patch('model_manager.activities.activities.Gates.__init__')
|
@patch('model_manager.activities.activities.Gates.__init__')
|
||||||
|
@patch('model_manager.activities.activities.Training.__init__')
|
||||||
def test___init__(
|
def test___init__(
|
||||||
mock_gates_init, mock_minio_init, mock_mlflow_init, mock_experiment_tracking_init
|
mock_training_init,
|
||||||
|
mock_gates_init,
|
||||||
|
mock_minio_init,
|
||||||
|
mock_mlflow_init,
|
||||||
|
mock_experiment_tracking_init,
|
||||||
):
|
):
|
||||||
postgres_config = {
|
postgres_config = {
|
||||||
'host': 'localhost',
|
'host': 'localhost',
|
||||||
@@ -54,6 +60,7 @@ def test___init__(
|
|||||||
assert isinstance(activities, ExperimentTracking)
|
assert isinstance(activities, ExperimentTracking)
|
||||||
assert isinstance(activities, MLFlow)
|
assert isinstance(activities, MLFlow)
|
||||||
assert isinstance(activities, Gates)
|
assert isinstance(activities, Gates)
|
||||||
|
assert isinstance(activities, Training)
|
||||||
|
|
||||||
mock_experiment_tracking_init.assert_called_once_with(
|
mock_experiment_tracking_init.assert_called_once_with(
|
||||||
ANY,
|
ANY,
|
||||||
@@ -97,6 +104,10 @@ def test___init__(
|
|||||||
ANY, logger=logger, notification_handler=notification_handler
|
ANY, logger=logger, notification_handler=notification_handler
|
||||||
)
|
)
|
||||||
|
|
||||||
|
mock_training_init.assert_called_once_with(
|
||||||
|
ANY, logger=logger, notification_handler=notification_handler
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@mark.asyncio
|
@mark.asyncio
|
||||||
@patch('model_manager.activities.activities.ExperimentTracking', return_value=MagicMock())
|
@patch('model_manager.activities.activities.ExperimentTracking', return_value=MagicMock())
|
||||||
|
|||||||
288
tests/activities/test_training.py
Normal file
288
tests/activities/test_training.py
Normal file
@@ -0,0 +1,288 @@
|
|||||||
|
"""Unit tests for Training activity."""
|
||||||
|
|
||||||
|
from io import BytesIO
|
||||||
|
from unittest.mock import MagicMock, patch
|
||||||
|
|
||||||
|
from pytest import mark
|
||||||
|
|
||||||
|
from model_manager.activities.training import Training
|
||||||
|
from model_manager.utils.models.train_model_result import TrainModelResult
|
||||||
|
|
||||||
|
|
||||||
|
@mark.asyncio
|
||||||
|
@patch('model_manager.activities.training.TrainingRepository')
|
||||||
|
async def test_train_model_success(mock_training_repository_class):
|
||||||
|
"""Test successful model training."""
|
||||||
|
# Create mock repository instance
|
||||||
|
mock_repository = MagicMock()
|
||||||
|
mock_training_repository_class.return_value = mock_repository
|
||||||
|
|
||||||
|
# Create mock train result
|
||||||
|
mock_train_result = MagicMock(spec=TrainModelResult)
|
||||||
|
mock_train_result.mse_val = 0.5
|
||||||
|
mock_train_result.mae_val = 0.3
|
||||||
|
mock_train_result.r2_val = 0.95
|
||||||
|
|
||||||
|
mock_final_result = MagicMock(spec=TrainModelResult)
|
||||||
|
mock_final_result.mse_val = 0.5
|
||||||
|
mock_final_result.mae_val = 0.3
|
||||||
|
mock_final_result.r2_val = 0.95
|
||||||
|
|
||||||
|
# Setup repository mocks
|
||||||
|
mock_repository.train.return_value = mock_train_result
|
||||||
|
mock_repository.after_train_calculation.return_value = mock_final_result
|
||||||
|
|
||||||
|
# Create Training instance
|
||||||
|
logger = MagicMock()
|
||||||
|
notification_handler = MagicMock()
|
||||||
|
training = Training(logger=logger, notification_handler=notification_handler)
|
||||||
|
|
||||||
|
# Mock inherited methods
|
||||||
|
training.info = MagicMock()
|
||||||
|
|
||||||
|
# Test data
|
||||||
|
uploaded_file = BytesIO(b'test,data\n1,2\n3,4')
|
||||||
|
train_params_dict = {
|
||||||
|
'experiment_run_id': 123,
|
||||||
|
'target_variable': 'price',
|
||||||
|
'variable_columns': ['feature1', 'feature2'],
|
||||||
|
'train_size': 80,
|
||||||
|
'shuffle': True,
|
||||||
|
'use_scaler': True,
|
||||||
|
'include_ar': False,
|
||||||
|
'bucket_name': 'test-bucket',
|
||||||
|
'file_name': 'test.csv',
|
||||||
|
'line_separator': '\n',
|
||||||
|
'decimal_separator': '.',
|
||||||
|
'lag_train': 1,
|
||||||
|
'lag_val': 1,
|
||||||
|
'rem_static_win': False,
|
||||||
|
'low_lim': {'feature1': 0.0, 'feature2': 0.0},
|
||||||
|
'upp_lim': {'feature1': 100.0, 'feature2': 100.0},
|
||||||
|
'window': 10,
|
||||||
|
'experiment_name': 'test_experiment',
|
||||||
|
'experiment_description': 'Test experiment',
|
||||||
|
'removed_intervals': [],
|
||||||
|
}
|
||||||
|
|
||||||
|
input_data = {
|
||||||
|
'metadata': {'workflow_id': 'test-123'},
|
||||||
|
'uploaded_file': uploaded_file,
|
||||||
|
'train_params': train_params_dict,
|
||||||
|
}
|
||||||
|
|
||||||
|
# Execute
|
||||||
|
result = await training.train_model(input_data)
|
||||||
|
|
||||||
|
# Assertions
|
||||||
|
assert result['success'] is True
|
||||||
|
assert result['result'] == mock_final_result
|
||||||
|
assert result['error_message'] is None
|
||||||
|
|
||||||
|
# Verify repository calls
|
||||||
|
mock_repository.train.assert_called_once()
|
||||||
|
mock_repository.after_train_calculation.assert_called_once()
|
||||||
|
|
||||||
|
|
||||||
|
@mark.asyncio
|
||||||
|
@patch('model_manager.activities.training.TrainingRepository')
|
||||||
|
async def test_train_model_invalid_file_type(mock_training_repository_class):
|
||||||
|
"""Test training with invalid file type."""
|
||||||
|
mock_repository = MagicMock()
|
||||||
|
mock_training_repository_class.return_value = mock_repository
|
||||||
|
|
||||||
|
logger = MagicMock()
|
||||||
|
notification_handler = MagicMock()
|
||||||
|
training = Training(logger=logger, notification_handler=notification_handler)
|
||||||
|
|
||||||
|
# Invalid file type (string instead of BytesIO)
|
||||||
|
input_data = {
|
||||||
|
'metadata': {},
|
||||||
|
'uploaded_file': 'not_a_bytesio',
|
||||||
|
'train_params': {
|
||||||
|
'experiment_run_id': 123,
|
||||||
|
'target_variable': 'price',
|
||||||
|
'variable_columns': ['feature1'],
|
||||||
|
'train_size': 80,
|
||||||
|
'shuffle': True,
|
||||||
|
'use_scaler': False,
|
||||||
|
'include_ar': False,
|
||||||
|
'bucket_name': 'test',
|
||||||
|
'file_name': 'test.csv',
|
||||||
|
'line_separator': '\n',
|
||||||
|
'decimal_separator': '.',
|
||||||
|
'lag_train': 1,
|
||||||
|
'lag_val': 1,
|
||||||
|
'rem_static_win': False,
|
||||||
|
'low_lim': {'feature1': 0.0},
|
||||||
|
'upp_lim': {'feature1': 100.0},
|
||||||
|
'window': 10,
|
||||||
|
'experiment_name': 'test_experiment',
|
||||||
|
'experiment_description': 'Test experiment',
|
||||||
|
'removed_intervals': [],
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result = await training.train_model(input_data)
|
||||||
|
|
||||||
|
assert result['success'] is False
|
||||||
|
assert result['result'] is None
|
||||||
|
assert 'uploaded_file must be BytesIO' in result['error_message']
|
||||||
|
# Verify notification was sent (via BaseActivity)
|
||||||
|
notification_handler.send_notification.assert_called_once()
|
||||||
|
|
||||||
|
|
||||||
|
@mark.asyncio
|
||||||
|
@patch('model_manager.activities.training.TrainingRepository')
|
||||||
|
async def test_train_model_training_error(mock_training_repository_class):
|
||||||
|
"""Test training failure during model training."""
|
||||||
|
mock_repository = MagicMock()
|
||||||
|
mock_training_repository_class.return_value = mock_repository
|
||||||
|
|
||||||
|
# Setup repository to raise error
|
||||||
|
mock_repository.train.side_effect = ValueError('Training data is empty')
|
||||||
|
|
||||||
|
logger = MagicMock()
|
||||||
|
notification_handler = MagicMock()
|
||||||
|
training = Training(logger=logger, notification_handler=notification_handler)
|
||||||
|
|
||||||
|
uploaded_file = BytesIO(b'test,data\n')
|
||||||
|
input_data = {
|
||||||
|
'metadata': {'workflow_id': 'test-456'},
|
||||||
|
'uploaded_file': uploaded_file,
|
||||||
|
'train_params': {
|
||||||
|
'experiment_run_id': 456,
|
||||||
|
'target_variable': 'price',
|
||||||
|
'variable_columns': ['feature1'],
|
||||||
|
'train_size': 80,
|
||||||
|
'shuffle': True,
|
||||||
|
'use_scaler': False,
|
||||||
|
'include_ar': False,
|
||||||
|
'bucket_name': 'test',
|
||||||
|
'file_name': 'test.csv',
|
||||||
|
'line_separator': '\n',
|
||||||
|
'decimal_separator': '.',
|
||||||
|
'lag_train': 1,
|
||||||
|
'lag_val': 1,
|
||||||
|
'rem_static_win': False,
|
||||||
|
'low_lim': {'feature1': 0.0},
|
||||||
|
'upp_lim': {'feature1': 100.0},
|
||||||
|
'window': 10,
|
||||||
|
'experiment_name': 'test_experiment',
|
||||||
|
'experiment_description': 'Test experiment',
|
||||||
|
'removed_intervals': [],
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result = await training.train_model(input_data)
|
||||||
|
|
||||||
|
assert result['success'] is False
|
||||||
|
assert result['result'] is None
|
||||||
|
assert 'Training data is empty' in result['error_message']
|
||||||
|
# Verify notification was sent (via BaseActivity)
|
||||||
|
notification_handler.send_notification.assert_called_once()
|
||||||
|
|
||||||
|
|
||||||
|
@mark.asyncio
|
||||||
|
@patch('model_manager.activities.training.TrainingRepository')
|
||||||
|
async def test_train_model_sends_notification_on_error(mock_training_repository_class):
|
||||||
|
"""Test that notification is sent when training fails."""
|
||||||
|
mock_repository = MagicMock()
|
||||||
|
mock_training_repository_class.return_value = mock_repository
|
||||||
|
|
||||||
|
mock_repository.train.side_effect = Exception('Database connection failed')
|
||||||
|
|
||||||
|
logger = MagicMock()
|
||||||
|
notification_handler = MagicMock()
|
||||||
|
training = Training(logger=logger, notification_handler=notification_handler)
|
||||||
|
|
||||||
|
uploaded_file = BytesIO(b'test,data\n1,2')
|
||||||
|
input_data = {
|
||||||
|
'metadata': {'workflow_id': 'test-789', 'experiment_run_id': 789},
|
||||||
|
'uploaded_file': uploaded_file,
|
||||||
|
'train_params': {
|
||||||
|
'experiment_run_id': 789,
|
||||||
|
'target_variable': 'price',
|
||||||
|
'variable_columns': ['feature1'],
|
||||||
|
'train_size': 80,
|
||||||
|
'shuffle': True,
|
||||||
|
'use_scaler': False,
|
||||||
|
'include_ar': False,
|
||||||
|
'bucket_name': 'test',
|
||||||
|
'file_name': 'test.csv',
|
||||||
|
'line_separator': '\n',
|
||||||
|
'decimal_separator': '.',
|
||||||
|
'lag_train': 1,
|
||||||
|
'lag_val': 1,
|
||||||
|
'rem_static_win': False,
|
||||||
|
'low_lim': {'feature1': 0.0},
|
||||||
|
'upp_lim': {'feature1': 100.0},
|
||||||
|
'window': 10,
|
||||||
|
'experiment_name': 'test_experiment',
|
||||||
|
'experiment_description': 'Test experiment',
|
||||||
|
'removed_intervals': [],
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result = await training.train_model(input_data)
|
||||||
|
|
||||||
|
# Verify notification was sent (via BaseActivity)
|
||||||
|
notification_handler.send_notification.assert_called_once()
|
||||||
|
|
||||||
|
# Verify result - the important part is that error was caught and returned
|
||||||
|
assert result['success'] is False
|
||||||
|
assert result['result'] is None
|
||||||
|
assert 'Database connection failed' in result['error_message']
|
||||||
|
|
||||||
|
|
||||||
|
@mark.asyncio
|
||||||
|
@patch('model_manager.activities.training.TrainingRepository')
|
||||||
|
async def test_train_model_after_calculation_error(mock_training_repository_class):
|
||||||
|
"""Test training failure during post-training calculations."""
|
||||||
|
mock_repository = MagicMock()
|
||||||
|
mock_training_repository_class.return_value = mock_repository
|
||||||
|
|
||||||
|
# Train succeeds but after_calculation fails
|
||||||
|
mock_train_result = MagicMock(spec=TrainModelResult)
|
||||||
|
mock_repository.train.return_value = mock_train_result
|
||||||
|
mock_repository.after_train_calculation.side_effect = Exception('Metric calculation failed')
|
||||||
|
|
||||||
|
logger = MagicMock()
|
||||||
|
notification_handler = MagicMock()
|
||||||
|
training = Training(logger=logger, notification_handler=notification_handler)
|
||||||
|
|
||||||
|
uploaded_file = BytesIO(b'test,data\n1,2\n3,4')
|
||||||
|
input_data = {
|
||||||
|
'metadata': {},
|
||||||
|
'uploaded_file': uploaded_file,
|
||||||
|
'train_params': {
|
||||||
|
'experiment_run_id': 999,
|
||||||
|
'target_variable': 'price',
|
||||||
|
'variable_columns': ['feature1'],
|
||||||
|
'train_size': 80,
|
||||||
|
'shuffle': True,
|
||||||
|
'use_scaler': False,
|
||||||
|
'include_ar': False,
|
||||||
|
'bucket_name': 'test',
|
||||||
|
'file_name': 'test.csv',
|
||||||
|
'line_separator': '\n',
|
||||||
|
'decimal_separator': '.',
|
||||||
|
'lag_train': 1,
|
||||||
|
'lag_val': 1,
|
||||||
|
'rem_static_win': False,
|
||||||
|
'low_lim': {'feature1': 0.0},
|
||||||
|
'upp_lim': {'feature1': 100.0},
|
||||||
|
'window': 10,
|
||||||
|
'experiment_name': 'test_experiment',
|
||||||
|
'experiment_description': 'Test experiment',
|
||||||
|
'removed_intervals': [],
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
result = await training.train_model(input_data)
|
||||||
|
|
||||||
|
assert result['success'] is False
|
||||||
|
assert result['result'] is None
|
||||||
|
assert 'Metric calculation failed' in result['error_message']
|
||||||
|
# Verify notification was sent (via BaseActivity)
|
||||||
|
notification_handler.send_notification.assert_called_once()
|
||||||
327
tests/utils/repository/test_training_repository.py
Normal file
327
tests/utils/repository/test_training_repository.py
Normal file
@@ -0,0 +1,327 @@
|
|||||||
|
"""Unit tests for TrainingRepository."""
|
||||||
|
|
||||||
|
from io import BytesIO
|
||||||
|
from unittest.mock import MagicMock, patch
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
import pandas as pd
|
||||||
|
from pytest import fixture, raises
|
||||||
|
|
||||||
|
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.training_repository import TrainingRepository
|
||||||
|
|
||||||
|
|
||||||
|
@fixture
|
||||||
|
def logger():
|
||||||
|
"""Create a mock logger."""
|
||||||
|
return MagicMock()
|
||||||
|
|
||||||
|
|
||||||
|
@fixture
|
||||||
|
def training_repository(logger):
|
||||||
|
"""Create a TrainingRepository instance."""
|
||||||
|
return TrainingRepository(logger)
|
||||||
|
|
||||||
|
|
||||||
|
@fixture
|
||||||
|
def train_params():
|
||||||
|
"""Create sample training parameters."""
|
||||||
|
return TrainModelParams(
|
||||||
|
variable_columns=['feature1', 'feature2'],
|
||||||
|
lag_train=1,
|
||||||
|
lag_val=1,
|
||||||
|
target_variable='target',
|
||||||
|
rem_static_win=False,
|
||||||
|
low_lim={'feature1': 0.0, 'feature2': 0.0},
|
||||||
|
upp_lim={'feature1': 100.0, 'feature2': 100.0},
|
||||||
|
window=10,
|
||||||
|
use_scaler=True,
|
||||||
|
include_ar=False,
|
||||||
|
bucket_name='test-bucket',
|
||||||
|
file_name='test.csv',
|
||||||
|
line_separator='\n',
|
||||||
|
decimal_separator='.',
|
||||||
|
train_size=80,
|
||||||
|
shuffle=True,
|
||||||
|
experiment_run_id=123,
|
||||||
|
experiment_name='test_experiment',
|
||||||
|
experiment_description='Test experiment',
|
||||||
|
removed_intervals=[],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@fixture
|
||||||
|
def sample_csv_data():
|
||||||
|
"""Create sample CSV data."""
|
||||||
|
csv_content = """feature1,feature2,target
|
||||||
|
1.0,2.0,10.0
|
||||||
|
2.0,3.0,15.0
|
||||||
|
3.0,4.0,20.0
|
||||||
|
4.0,5.0,25.0
|
||||||
|
5.0,6.0,30.0
|
||||||
|
6.0,7.0,35.0
|
||||||
|
7.0,8.0,40.0
|
||||||
|
8.0,9.0,45.0
|
||||||
|
9.0,10.0,50.0
|
||||||
|
10.0,11.0,55.0
|
||||||
|
"""
|
||||||
|
return BytesIO(csv_content.encode())
|
||||||
|
|
||||||
|
|
||||||
|
@patch('model_manager.utils.repository.training_repository.load_data')
|
||||||
|
@patch('model_manager.utils.repository.training_repository.DataPreprocessor')
|
||||||
|
@patch('model_manager.utils.repository.training_repository.split_train_test')
|
||||||
|
@patch('model_manager.utils.repository.training_repository.LinearRegressionModel')
|
||||||
|
def test_train_success(
|
||||||
|
mock_linear_model,
|
||||||
|
mock_split,
|
||||||
|
mock_preprocessor_class,
|
||||||
|
mock_load_data,
|
||||||
|
training_repository,
|
||||||
|
train_params,
|
||||||
|
sample_csv_data,
|
||||||
|
):
|
||||||
|
"""Test successful model training."""
|
||||||
|
# Setup mocks
|
||||||
|
mock_data = pd.DataFrame(
|
||||||
|
{'feature1': [1, 2, 3, 4, 5], 'feature2': [2, 3, 4, 5, 6], 'target': [10, 15, 20, 25, 30]}
|
||||||
|
)
|
||||||
|
mock_load_data.return_value = mock_data
|
||||||
|
|
||||||
|
mock_preprocessor = MagicMock()
|
||||||
|
mock_preprocessor_class.return_value = mock_preprocessor
|
||||||
|
mock_preprocessor.transform.return_value = mock_data
|
||||||
|
|
||||||
|
x_train = pd.DataFrame({'feature1': [1, 2, 3], 'feature2': [2, 3, 4]})
|
||||||
|
x_test = pd.DataFrame({'feature1': [4, 5], 'feature2': [5, 6]})
|
||||||
|
y_train = pd.Series([10, 15, 20], name='target')
|
||||||
|
y_test = pd.Series([25, 30], name='target')
|
||||||
|
mock_split.return_value = (x_train, x_test, y_train, y_test)
|
||||||
|
|
||||||
|
mock_model = MagicMock()
|
||||||
|
mock_linear_model.return_value = mock_model
|
||||||
|
|
||||||
|
mock_scaler = MagicMock()
|
||||||
|
mock_preprocessor.get_scaler.return_value = mock_scaler
|
||||||
|
|
||||||
|
# Execute
|
||||||
|
result = training_repository.train(sample_csv_data, train_params)
|
||||||
|
|
||||||
|
# Assertions
|
||||||
|
assert isinstance(result, TrainModelResult)
|
||||||
|
assert result.params == train_params
|
||||||
|
assert result.process_data == mock_preprocessor
|
||||||
|
assert result.regr == mock_model
|
||||||
|
mock_load_data.assert_called_once()
|
||||||
|
mock_preprocessor.fit.assert_called_once()
|
||||||
|
mock_model.fit.assert_called_once()
|
||||||
|
|
||||||
|
|
||||||
|
@patch('model_manager.utils.repository.training_repository.load_data')
|
||||||
|
def test_train_empty_data_after_transform(
|
||||||
|
mock_load_data, training_repository, train_params, sample_csv_data
|
||||||
|
):
|
||||||
|
"""Test training with empty data after transformation."""
|
||||||
|
mock_data = pd.DataFrame({'feature1': [], 'feature2': [], 'target': []})
|
||||||
|
mock_load_data.return_value = mock_data
|
||||||
|
|
||||||
|
with patch.object(training_repository, 'init_data_preprocessor') as mock_init:
|
||||||
|
mock_preprocessor = MagicMock()
|
||||||
|
mock_init.return_value = mock_preprocessor
|
||||||
|
mock_preprocessor.transform.return_value = pd.DataFrame()
|
||||||
|
|
||||||
|
with raises(ValueError, match='Data view is empty after transformation'):
|
||||||
|
training_repository.train(sample_csv_data, train_params)
|
||||||
|
|
||||||
|
|
||||||
|
def test_init_scaler_dict_with_minmax_scaler(training_repository, train_params):
|
||||||
|
"""Test scaler dict initialization with MinMaxScaler."""
|
||||||
|
mock_preprocessor = MagicMock()
|
||||||
|
mock_scaler = MagicMock()
|
||||||
|
mock_scaler.x_min = [0.0, 1.0]
|
||||||
|
mock_scaler.x_max = [10.0, 11.0]
|
||||||
|
mock_scaler.y_min = 5.0
|
||||||
|
mock_scaler.y_max = 50.0
|
||||||
|
mock_preprocessor.get_scaler.return_value = mock_scaler
|
||||||
|
|
||||||
|
# Patch isinstance to return True for MinMaxScaler
|
||||||
|
with patch(
|
||||||
|
'model_manager.utils.repository.training_repository.isinstance',
|
||||||
|
side_effect=lambda obj, cls: cls.__name__ == 'MinMaxScaler',
|
||||||
|
):
|
||||||
|
result = training_repository.init_scaler_dict(mock_preprocessor, train_params)
|
||||||
|
|
||||||
|
assert result is not None
|
||||||
|
assert 'feature1' in result
|
||||||
|
assert 'feature2' in result
|
||||||
|
assert 'target' in result
|
||||||
|
assert result['feature1'] == {'min': 0.0, 'max': 10.0}
|
||||||
|
assert result['feature2'] == {'min': 1.0, 'max': 11.0}
|
||||||
|
assert result['target'] == {'min': 5.0, 'max': 50.0}
|
||||||
|
|
||||||
|
|
||||||
|
def test_init_scaler_dict_with_z_scaler(training_repository, train_params):
|
||||||
|
"""Test scaler dict initialization with Z_Scaler."""
|
||||||
|
mock_preprocessor = MagicMock()
|
||||||
|
mock_scaler = MagicMock()
|
||||||
|
mock_scaler.create_dict.return_value = {'mean': 5.0, 'std': 2.0}
|
||||||
|
mock_preprocessor.get_scaler.return_value = mock_scaler
|
||||||
|
|
||||||
|
# Patch isinstance to return True for Z_Scaler
|
||||||
|
with patch(
|
||||||
|
'model_manager.utils.repository.training_repository.isinstance',
|
||||||
|
side_effect=lambda obj, cls: cls.__name__ == 'Z_Scaler',
|
||||||
|
):
|
||||||
|
result = training_repository.init_scaler_dict(mock_preprocessor, train_params)
|
||||||
|
|
||||||
|
assert result == {'mean': 5.0, 'std': 2.0}
|
||||||
|
mock_scaler.create_dict.assert_called_once()
|
||||||
|
|
||||||
|
|
||||||
|
def test_init_scaler_dict_without_scaler(training_repository):
|
||||||
|
"""Test scaler dict initialization when use_scaler is False."""
|
||||||
|
train_params_no_scaler = TrainModelParams(
|
||||||
|
variable_columns=['feature1'],
|
||||||
|
lag_train=1,
|
||||||
|
lag_val=1,
|
||||||
|
target_variable='target',
|
||||||
|
rem_static_win=False,
|
||||||
|
low_lim={'feature1': 0.0},
|
||||||
|
upp_lim={'feature1': 100.0},
|
||||||
|
window=10,
|
||||||
|
use_scaler=False,
|
||||||
|
include_ar=False,
|
||||||
|
bucket_name='test',
|
||||||
|
file_name='test.csv',
|
||||||
|
line_separator='\n',
|
||||||
|
decimal_separator='.',
|
||||||
|
train_size=80,
|
||||||
|
shuffle=True,
|
||||||
|
experiment_run_id=123,
|
||||||
|
experiment_name='test',
|
||||||
|
experiment_description='test',
|
||||||
|
removed_intervals=[],
|
||||||
|
)
|
||||||
|
|
||||||
|
mock_preprocessor = MagicMock()
|
||||||
|
result = training_repository.init_scaler_dict(mock_preprocessor, train_params_no_scaler)
|
||||||
|
|
||||||
|
assert result == {}
|
||||||
|
|
||||||
|
|
||||||
|
@patch('model_manager.utils.repository.training_repository.mse')
|
||||||
|
@patch('model_manager.utils.repository.training_repository.mae')
|
||||||
|
@patch('model_manager.utils.repository.training_repository.r2')
|
||||||
|
def test_after_train_calculation_with_scaler(
|
||||||
|
mock_r2, mock_mae, mock_mse, training_repository, train_params
|
||||||
|
):
|
||||||
|
"""Test post-training calculations with scaler."""
|
||||||
|
# Setup mock train result
|
||||||
|
mock_train_result = MagicMock(spec=TrainModelResult)
|
||||||
|
mock_train_result.params = train_params
|
||||||
|
mock_train_result.x_train = pd.DataFrame({'feature1': [1, 2, 3], 'feature2': [2, 3, 4]})
|
||||||
|
mock_train_result.x_test = pd.DataFrame({'feature1': [4, 5], 'feature2': [5, 6]})
|
||||||
|
mock_train_result.y_train = pd.Series([10, 15, 20], name='target')
|
||||||
|
mock_train_result.y_test = pd.Series([25, 30], name='target')
|
||||||
|
|
||||||
|
mock_regr = MagicMock()
|
||||||
|
mock_regr.predict.return_value = np.array([24.5, 29.5])
|
||||||
|
mock_train_result.regr = mock_regr
|
||||||
|
|
||||||
|
mock_scaler = MagicMock()
|
||||||
|
mock_scaler.denormalize_single_input.side_effect = lambda x, col: x
|
||||||
|
mock_scaler.denormalize_predictions.side_effect = lambda x, col: x
|
||||||
|
|
||||||
|
mock_process_data = MagicMock()
|
||||||
|
mock_process_data.get_scaler.return_value = mock_scaler
|
||||||
|
mock_train_result.process_data = mock_process_data
|
||||||
|
|
||||||
|
# Setup metric mocks
|
||||||
|
mock_mse.return_value = 0.5
|
||||||
|
mock_mae.return_value = 0.3
|
||||||
|
mock_r2.return_value = 0.95
|
||||||
|
|
||||||
|
# Execute
|
||||||
|
result = training_repository.after_train_calculation(train_params, mock_train_result)
|
||||||
|
|
||||||
|
# Assertions
|
||||||
|
assert result == mock_train_result
|
||||||
|
assert result.mse_val == 0.5
|
||||||
|
assert result.mae_val == 0.3
|
||||||
|
assert result.r2_val == 0.95
|
||||||
|
assert result.y_pred is not None
|
||||||
|
mock_regr.predict.assert_called_once()
|
||||||
|
mock_mse.assert_called_once()
|
||||||
|
mock_mae.assert_called_once()
|
||||||
|
mock_r2.assert_called_once()
|
||||||
|
|
||||||
|
|
||||||
|
@patch('model_manager.utils.repository.training_repository.mse')
|
||||||
|
@patch('model_manager.utils.repository.training_repository.mae')
|
||||||
|
@patch('model_manager.utils.repository.training_repository.r2')
|
||||||
|
def test_after_train_calculation_without_scaler(mock_r2, mock_mae, mock_mse, training_repository):
|
||||||
|
"""Test post-training calculations without scaler."""
|
||||||
|
train_params_no_scaler = TrainModelParams(
|
||||||
|
variable_columns=['feature1'],
|
||||||
|
lag_train=1,
|
||||||
|
lag_val=1,
|
||||||
|
target_variable='target',
|
||||||
|
rem_static_win=False,
|
||||||
|
low_lim={'feature1': 0.0},
|
||||||
|
upp_lim={'feature1': 100.0},
|
||||||
|
window=10,
|
||||||
|
use_scaler=False,
|
||||||
|
include_ar=False,
|
||||||
|
bucket_name='test',
|
||||||
|
file_name='test.csv',
|
||||||
|
line_separator='\n',
|
||||||
|
decimal_separator='.',
|
||||||
|
train_size=80,
|
||||||
|
shuffle=True,
|
||||||
|
experiment_run_id=123,
|
||||||
|
experiment_name='test',
|
||||||
|
experiment_description='test',
|
||||||
|
removed_intervals=[],
|
||||||
|
)
|
||||||
|
|
||||||
|
mock_train_result = MagicMock(spec=TrainModelResult)
|
||||||
|
mock_train_result.params = train_params_no_scaler
|
||||||
|
mock_train_result.x_train = pd.DataFrame({'feature1': [1, 2, 3]})
|
||||||
|
mock_train_result.x_test = pd.DataFrame({'feature1': [4, 5]})
|
||||||
|
mock_train_result.y_train = pd.Series([10, 15, 20], name='target')
|
||||||
|
mock_train_result.y_test = pd.Series([25, 30], name='target')
|
||||||
|
|
||||||
|
mock_regr = MagicMock()
|
||||||
|
mock_regr.predict.return_value = np.array([24.5, 29.5])
|
||||||
|
mock_train_result.regr = mock_regr
|
||||||
|
|
||||||
|
# Setup metric mocks
|
||||||
|
mock_mse.return_value = 0.5
|
||||||
|
mock_mae.return_value = 0.3
|
||||||
|
mock_r2.return_value = 0.95
|
||||||
|
|
||||||
|
# Execute
|
||||||
|
result = training_repository.after_train_calculation(train_params_no_scaler, mock_train_result)
|
||||||
|
|
||||||
|
# Assertions
|
||||||
|
assert result.mse_val == 0.5
|
||||||
|
assert result.mae_val == 0.3
|
||||||
|
assert result.r2_val == 0.95
|
||||||
|
|
||||||
|
|
||||||
|
@patch('model_manager.utils.repository.training_repository.DataPreprocessor')
|
||||||
|
def test_init_data_preprocessor(mock_preprocessor_class, training_repository, train_params):
|
||||||
|
"""Test DataPreprocessor initialization."""
|
||||||
|
mock_preprocessor = MagicMock()
|
||||||
|
mock_preprocessor_class.return_value = mock_preprocessor
|
||||||
|
|
||||||
|
result = training_repository.init_data_preprocessor(train_params)
|
||||||
|
|
||||||
|
assert result == mock_preprocessor
|
||||||
|
mock_preprocessor_class.assert_called_once()
|
||||||
|
call_kwargs = mock_preprocessor_class.call_args[1]
|
||||||
|
assert call_kwargs['target_variable'] == 'target'
|
||||||
|
assert call_kwargs['input_columns'] == ['feature1', 'feature2']
|
||||||
|
assert call_kwargs['low_lim'] == {'feature1': 0.0, 'feature2': 0.0}
|
||||||
|
assert call_kwargs['upp_lim'] == {'feature1': 100.0, 'feature2': 100.0}
|
||||||
Reference in New Issue
Block a user