SIENTIAPDE-1251: Implement ML model training activity and repository
This commit introduces the 'Training' activity and 'TrainingRepository' for handling ML model training operations within the Model Manager system. - Added model_manager/activities/training.py for the Training activity, which extends BaseActivity and integrates with Temporal workflows. - Added model_manager/utils/repository/training_repository.py for the TrainingRepository, which encapsulates the core training logic. - Updated model_manager/activities/activities.py to include the Training activity in the main activities orchestrator. - Updated README.md to document the new 'Training' component. - Added unit tests for the new activity and repository.
This commit is contained in:
@@ -6,14 +6,20 @@ from model_manager.activities.activities import Activities
|
||||
from model_manager.activities.experiment_tracking import ExperimentTracking
|
||||
from model_manager.activities.gates import Gates
|
||||
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.MLFlow.__init__')
|
||||
@patch('model_manager.activities.activities.MinIO.__init__')
|
||||
@patch('model_manager.activities.activities.Gates.__init__')
|
||||
@patch('model_manager.activities.activities.Training.__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 = {
|
||||
'host': 'localhost',
|
||||
@@ -54,6 +60,7 @@ def test___init__(
|
||||
assert isinstance(activities, ExperimentTracking)
|
||||
assert isinstance(activities, MLFlow)
|
||||
assert isinstance(activities, Gates)
|
||||
assert isinstance(activities, Training)
|
||||
|
||||
mock_experiment_tracking_init.assert_called_once_with(
|
||||
ANY,
|
||||
@@ -97,6 +104,10 @@ def test___init__(
|
||||
ANY, logger=logger, notification_handler=notification_handler
|
||||
)
|
||||
|
||||
mock_training_init.assert_called_once_with(
|
||||
ANY, logger=logger, notification_handler=notification_handler
|
||||
)
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@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