SIENTIAPDE-1579: Updated tests
This commit is contained in:
@@ -344,9 +344,14 @@ def test_create_run_directory_success(
|
||||
|
||||
result = repo._create_run_directory('/tmp/reports', 'test_run') # noqa: S108
|
||||
|
||||
expected_path = os.path.join('/tmp/reports/temp', 'test_run_20240101_120000_123456') # noqa: S108
|
||||
assert result == expected_path
|
||||
mock_makedirs.assert_called_once_with(expected_path, exist_ok=True)
|
||||
expected_path = os.path.normpath(
|
||||
os.path.join('/tmp/reports', 'temp', 'test_run_20240101_120000_123456') # noqa: S108
|
||||
)
|
||||
assert os.path.normpath(result) == expected_path
|
||||
mock_makedirs.assert_called_once()
|
||||
call_path = mock_makedirs.call_args[0][0]
|
||||
assert os.path.normpath(call_path) == expected_path
|
||||
assert mock_makedirs.call_args[1] == {'exist_ok': True}
|
||||
|
||||
|
||||
@patch('model_manager.utils.repository.model_repository.ModelServing')
|
||||
|
||||
@@ -10,7 +10,10 @@ import pytest
|
||||
from model_manager.sientia.models import LinearRegressionModel
|
||||
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
|
||||
from model_manager.utils.repository.training_repository import (
|
||||
TrainingRepository,
|
||||
_apply_support_filters,
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
@@ -57,6 +60,8 @@ def sample_params():
|
||||
scaler_name='None',
|
||||
support_filters={},
|
||||
static_threshold=None,
|
||||
date_column=None,
|
||||
date_format=None,
|
||||
)
|
||||
|
||||
|
||||
@@ -139,6 +144,8 @@ class TestExtractModelEquation:
|
||||
scaler_name='None',
|
||||
support_filters={},
|
||||
static_threshold=None,
|
||||
date_column=None,
|
||||
date_format=None,
|
||||
)
|
||||
|
||||
# Mock model with single coefficient
|
||||
@@ -585,7 +592,7 @@ class TestTrain:
|
||||
"""
|
||||
return BytesIO(csv_content.encode('utf-8'))
|
||||
|
||||
@patch.object(TrainingRepository, '_configure_datetime_index', lambda self, df: df)
|
||||
@patch.object(TrainingRepository, '_configure_datetime_index', lambda self, df, params: df)
|
||||
@patch('model_manager.utils.repository.training_repository.split_train_test')
|
||||
@patch('model_manager.utils.repository.training_repository.load_data')
|
||||
def test_train_basic_workflow(
|
||||
@@ -631,7 +638,7 @@ class TestTrain:
|
||||
# Verify split was called
|
||||
assert mock_split_train_test.called
|
||||
|
||||
@patch.object(TrainingRepository, '_configure_datetime_index', lambda self, df: df)
|
||||
@patch.object(TrainingRepository, '_configure_datetime_index', lambda self, df, params: df)
|
||||
@patch('model_manager.utils.repository.training_repository.split_train_test')
|
||||
@patch('model_manager.utils.repository.training_repository.load_data')
|
||||
def test_train_with_scaler(
|
||||
@@ -661,7 +668,7 @@ class TestTrain:
|
||||
assert result is not None
|
||||
assert result.scaler_dict is not None
|
||||
|
||||
@patch.object(TrainingRepository, '_configure_datetime_index', lambda self, df: df)
|
||||
@patch.object(TrainingRepository, '_configure_datetime_index', lambda self, df, params: df)
|
||||
@patch('model_manager.utils.repository.training_repository.split_train_test')
|
||||
@patch('model_manager.utils.repository.training_repository.load_data')
|
||||
def test_train_with_shuffle_enabled(
|
||||
@@ -692,7 +699,7 @@ class TestTrain:
|
||||
call_kwargs = mock_split_train_test.call_args[1]
|
||||
assert call_kwargs['shuffle'] is True
|
||||
|
||||
@patch.object(TrainingRepository, '_configure_datetime_index', lambda self, df: df)
|
||||
@patch.object(TrainingRepository, '_configure_datetime_index', lambda self, df, params: df)
|
||||
@patch('model_manager.utils.repository.training_repository.split_train_test')
|
||||
@patch('model_manager.utils.repository.training_repository.load_data')
|
||||
def test_train_with_different_train_size(
|
||||
@@ -723,7 +730,7 @@ class TestTrain:
|
||||
call_kwargs = mock_split_train_test.call_args[1]
|
||||
assert call_kwargs['train_size'] == 0.7
|
||||
|
||||
@patch.object(TrainingRepository, '_configure_datetime_index', lambda self, df: df)
|
||||
@patch.object(TrainingRepository, '_configure_datetime_index', lambda self, df, params: df)
|
||||
@patch('model_manager.utils.repository.training_repository.split_train_test')
|
||||
@patch('model_manager.utils.repository.training_repository.load_data')
|
||||
def test_train_raises_on_empty_data_after_transform(
|
||||
@@ -744,7 +751,7 @@ class TestTrain:
|
||||
with pytest.raises(ValueError, match='Data view is empty after transformation'):
|
||||
training_repo.train(sample_csv_data, sample_params)
|
||||
|
||||
@patch.object(TrainingRepository, '_configure_datetime_index', lambda self, df: df)
|
||||
@patch.object(TrainingRepository, '_configure_datetime_index', lambda self, df, params: df)
|
||||
@patch('model_manager.utils.repository.training_repository.split_train_test')
|
||||
@patch('model_manager.utils.repository.training_repository.load_data')
|
||||
def test_train_logs_success(
|
||||
@@ -781,7 +788,7 @@ class TestTrain:
|
||||
'Model trained successfully' in str(call) for call in mock_logger.info.call_args_list
|
||||
)
|
||||
|
||||
@patch.object(TrainingRepository, '_configure_datetime_index', lambda self, df: df)
|
||||
@patch.object(TrainingRepository, '_configure_datetime_index', lambda self, df, params: df)
|
||||
@patch('model_manager.utils.repository.training_repository.split_train_test')
|
||||
@patch('model_manager.utils.repository.training_repository.load_data')
|
||||
def test_train_with_custom_separators(
|
||||
@@ -812,7 +819,7 @@ class TestTrain:
|
||||
# Verify load_data was called with custom separators
|
||||
mock_load_data.assert_called_once_with(sample_csv_data, ';', ',')
|
||||
|
||||
@patch.object(TrainingRepository, '_configure_datetime_index', lambda self, df: df)
|
||||
@patch.object(TrainingRepository, '_configure_datetime_index', lambda self, df, params: df)
|
||||
@patch('model_manager.utils.repository.training_repository.split_train_test')
|
||||
@patch('model_manager.utils.repository.training_repository.load_data')
|
||||
def test_train_result_contains_all_fields(
|
||||
@@ -862,19 +869,55 @@ class TestConfigureDatetimeIndex:
|
||||
"""Create a TrainingRepository instance."""
|
||||
return TrainingRepository(logger=mock_logger)
|
||||
|
||||
def test_configure_datetime_index_already_datetime(self, training_repo):
|
||||
@pytest.fixture
|
||||
def datetime_params(self):
|
||||
"""Minimal params for _configure_datetime_index (date_column/date_format can be None)."""
|
||||
return TrainModelParams(
|
||||
experiment_run_id=1,
|
||||
experiment_name='test',
|
||||
target_variable='y',
|
||||
variable_columns=['x'],
|
||||
lag_train={'x': 0},
|
||||
lag_val={'x': 0},
|
||||
rem_static_win=False,
|
||||
low_lim={},
|
||||
upp_lim={},
|
||||
window=0,
|
||||
use_scaler=False,
|
||||
include_ar=False,
|
||||
train_size=80,
|
||||
shuffle=True,
|
||||
bucket_name='b',
|
||||
file_name='f.csv',
|
||||
line_separator=',',
|
||||
decimal_separator='.',
|
||||
removed_intervals=[],
|
||||
model_name='Linear Regression',
|
||||
degree=1,
|
||||
interaction_only=False,
|
||||
nan_treatment='drop',
|
||||
start_date=None,
|
||||
end_date=None,
|
||||
scaler_name='None',
|
||||
support_filters={},
|
||||
static_threshold=None,
|
||||
date_column=None,
|
||||
date_format=None,
|
||||
)
|
||||
|
||||
def test_configure_datetime_index_already_datetime(self, training_repo, datetime_params):
|
||||
"""Test _configure_datetime_index when index is already DatetimeIndex."""
|
||||
data = pd.DataFrame(
|
||||
{'var1': [1, 2, 3], 'var2': [4, 5, 6]},
|
||||
index=pd.to_datetime(['2023-01-01', '2023-01-02', '2023-01-03']),
|
||||
)
|
||||
|
||||
result = training_repo._configure_datetime_index(data)
|
||||
result = training_repo._configure_datetime_index(data, datetime_params)
|
||||
|
||||
assert isinstance(result.index, pd.DatetimeIndex)
|
||||
assert len(result) == 3
|
||||
|
||||
def test_configure_datetime_index_with_timestamp_column(self, training_repo):
|
||||
def test_configure_datetime_index_with_timestamp_column(self, training_repo, datetime_params):
|
||||
"""Test _configure_datetime_index with 'timestamp' column."""
|
||||
data = pd.DataFrame(
|
||||
{
|
||||
@@ -884,12 +927,12 @@ class TestConfigureDatetimeIndex:
|
||||
}
|
||||
)
|
||||
|
||||
result = training_repo._configure_datetime_index(data)
|
||||
result = training_repo._configure_datetime_index(data, datetime_params)
|
||||
|
||||
assert isinstance(result.index, pd.DatetimeIndex)
|
||||
assert 'timestamp' not in result.columns
|
||||
|
||||
def test_configure_datetime_index_with_date_column(self, training_repo):
|
||||
def test_configure_datetime_index_with_date_column(self, training_repo, datetime_params):
|
||||
"""Test _configure_datetime_index with 'date' column."""
|
||||
data = pd.DataFrame(
|
||||
{
|
||||
@@ -898,12 +941,12 @@ class TestConfigureDatetimeIndex:
|
||||
}
|
||||
)
|
||||
|
||||
result = training_repo._configure_datetime_index(data)
|
||||
result = training_repo._configure_datetime_index(data, datetime_params)
|
||||
|
||||
assert isinstance(result.index, pd.DatetimeIndex)
|
||||
assert 'date' not in result.columns
|
||||
|
||||
def test_configure_datetime_index_with_datetime_column(self, training_repo):
|
||||
def test_configure_datetime_index_with_datetime_column(self, training_repo, datetime_params):
|
||||
"""Test _configure_datetime_index with 'datetime' column."""
|
||||
data = pd.DataFrame(
|
||||
{
|
||||
@@ -912,12 +955,12 @@ class TestConfigureDatetimeIndex:
|
||||
}
|
||||
)
|
||||
|
||||
result = training_repo._configure_datetime_index(data)
|
||||
result = training_repo._configure_datetime_index(data, datetime_params)
|
||||
|
||||
assert isinstance(result.index, pd.DatetimeIndex)
|
||||
assert 'datetime' not in result.columns
|
||||
|
||||
def test_configure_datetime_index_first_column_datetime(self, training_repo):
|
||||
def test_configure_datetime_index_first_column_datetime(self, training_repo, datetime_params):
|
||||
"""Test _configure_datetime_index when first column looks like datetime."""
|
||||
data = pd.DataFrame(
|
||||
{
|
||||
@@ -926,13 +969,13 @@ class TestConfigureDatetimeIndex:
|
||||
}
|
||||
)
|
||||
|
||||
result = training_repo._configure_datetime_index(data)
|
||||
result = training_repo._configure_datetime_index(data, datetime_params)
|
||||
|
||||
assert isinstance(result.index, pd.DatetimeIndex)
|
||||
assert 'my_date' not in result.columns
|
||||
|
||||
@pytest.mark.filterwarnings('ignore::UserWarning')
|
||||
def test_configure_datetime_index_no_timestamp_column(self, training_repo):
|
||||
def test_configure_datetime_index_no_timestamp_column(self, training_repo, datetime_params):
|
||||
"""Test _configure_datetime_index when no timestamp column found."""
|
||||
data = pd.DataFrame(
|
||||
{
|
||||
@@ -941,14 +984,14 @@ class TestConfigureDatetimeIndex:
|
||||
}
|
||||
)
|
||||
|
||||
result = training_repo._configure_datetime_index(data)
|
||||
result = training_repo._configure_datetime_index(data, datetime_params)
|
||||
|
||||
# Should return original data unchanged (no valid datetime columns)
|
||||
assert 'var1' in result.columns
|
||||
assert 'var2' in result.columns
|
||||
|
||||
@pytest.mark.filterwarnings('ignore::UserWarning')
|
||||
def test_configure_datetime_index_invalid_timestamp_column(self, training_repo):
|
||||
def test_configure_datetime_index_invalid_timestamp_column(self, training_repo, datetime_params):
|
||||
"""Test _configure_datetime_index with invalid timestamp values."""
|
||||
data = pd.DataFrame(
|
||||
{
|
||||
@@ -957,13 +1000,13 @@ class TestConfigureDatetimeIndex:
|
||||
}
|
||||
)
|
||||
|
||||
result = training_repo._configure_datetime_index(data)
|
||||
result = training_repo._configure_datetime_index(data, datetime_params)
|
||||
|
||||
# Should skip invalid column and try first column
|
||||
assert 'var1' in result.columns
|
||||
|
||||
@pytest.mark.filterwarnings('ignore::UserWarning')
|
||||
def test_configure_datetime_index_invalid_first_column(self, training_repo):
|
||||
def test_configure_datetime_index_invalid_first_column(self, training_repo, datetime_params):
|
||||
"""Test _configure_datetime_index when first column is not datetime."""
|
||||
data = pd.DataFrame(
|
||||
{
|
||||
@@ -972,13 +1015,13 @@ class TestConfigureDatetimeIndex:
|
||||
}
|
||||
)
|
||||
|
||||
result = training_repo._configure_datetime_index(data)
|
||||
result = training_repo._configure_datetime_index(data, datetime_params)
|
||||
|
||||
# Should return original data unchanged
|
||||
assert 'var1' in result.columns
|
||||
assert 'var2' in result.columns
|
||||
|
||||
def test_configure_datetime_index_first_column_all_nan(self, training_repo):
|
||||
def test_configure_datetime_index_first_column_all_nan(self, training_repo, datetime_params):
|
||||
"""Test _configure_datetime_index when first column has all NaN values."""
|
||||
data = pd.DataFrame(
|
||||
{
|
||||
@@ -987,7 +1030,7 @@ class TestConfigureDatetimeIndex:
|
||||
}
|
||||
)
|
||||
|
||||
result = training_repo._configure_datetime_index(data)
|
||||
result = training_repo._configure_datetime_index(data, datetime_params)
|
||||
|
||||
# Should return original data unchanged (first column has no valid values)
|
||||
assert 'first_col' in result.columns
|
||||
@@ -1072,3 +1115,98 @@ class TestExtractModelEquationPolynomial:
|
||||
|
||||
assert 'equation_string' in result
|
||||
assert 'latex_equation' in result
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Tests for _apply_support_filters
|
||||
# ============================================================================
|
||||
|
||||
|
||||
class TestApplySupportFilters:
|
||||
"""Tests for _apply_support_filters function."""
|
||||
|
||||
def test_empty_support_filters_returns_unchanged(self):
|
||||
"""When support_filters is empty, data_view is returned unchanged."""
|
||||
data = pd.DataFrame({'x': [1, 2, 3], 'target': [10, 20, 30]})
|
||||
result = _apply_support_filters(data, 'target', {})
|
||||
pd.testing.assert_frame_equal(result, data)
|
||||
|
||||
def test_target_not_in_columns_returns_unchanged(self):
|
||||
"""When target_variable is not in data_view columns, return unchanged."""
|
||||
data = pd.DataFrame({'x': [1, 2, 3], 'y': [10, 20, 30]})
|
||||
result = _apply_support_filters(data, 'target', {'x': {}})
|
||||
pd.testing.assert_frame_equal(result, data)
|
||||
|
||||
def test_variable_not_in_columns_skipped(self):
|
||||
"""When a filter variable is not in data_view, that variable is skipped."""
|
||||
data = pd.DataFrame({'x': [1, 2, 3], 'target': [10, 20, 30]})
|
||||
support_filters = {
|
||||
'missing_var': {
|
||||
'upper_line': {'intercept': 100, 'angle': 10},
|
||||
'lower_line': {'intercept': 0, 'angle': -10},
|
||||
},
|
||||
}
|
||||
result = _apply_support_filters(data, 'target', support_filters)
|
||||
pd.testing.assert_frame_equal(result, data)
|
||||
|
||||
def test_snake_case_upper_lower_line(self):
|
||||
"""Support filters with upper_line/lower_line (snake_case) filter rows."""
|
||||
data = pd.DataFrame({
|
||||
'x': [1.0, 2.0, 3.0, 4.0],
|
||||
'target': [2.0, 4.0, 6.0, 8.0],
|
||||
})
|
||||
support_filters = {
|
||||
'x': {
|
||||
'upper_line': {'intercept': 1.0, 'angle': 50},
|
||||
'lower_line': {'intercept': -1.0, 'angle': -50},
|
||||
},
|
||||
}
|
||||
result = _apply_support_filters(data, 'target', support_filters)
|
||||
assert len(result) <= 4
|
||||
assert list(result.columns) == ['x', 'target']
|
||||
|
||||
def test_camel_case_upper_lower_line(self):
|
||||
"""Support filters with upperLine/lowerLine (camelCase) are accepted."""
|
||||
data = pd.DataFrame({
|
||||
'x': [1.0, 2.0, 3.0],
|
||||
'target': [1.0, 2.0, 3.0],
|
||||
})
|
||||
support_filters = {
|
||||
'x': {
|
||||
'upperLine': {'intercept': 2, 'angle': 5},
|
||||
'lowerLine': {'intercept': 0, 'angle': -5},
|
||||
},
|
||||
}
|
||||
result = _apply_support_filters(data, 'target', support_filters)
|
||||
assert len(result) <= 3
|
||||
assert list(result.columns) == ['x', 'target']
|
||||
|
||||
def test_two_variables_ands_masks(self):
|
||||
"""Two variables apply AND of both masks."""
|
||||
data = pd.DataFrame({
|
||||
'a': [1.0, 2.0, 3.0],
|
||||
'b': [1.0, 2.0, 3.0],
|
||||
'target': [2.0, 2.0, 2.0],
|
||||
})
|
||||
support_filters = {
|
||||
'a': {
|
||||
'upper_line': {'intercept': 10, 'angle': 45},
|
||||
'lower_line': {'intercept': -10, 'angle': -45},
|
||||
},
|
||||
'b': {
|
||||
'upper_line': {'intercept': 10, 'angle': 45},
|
||||
'lower_line': {'intercept': -10, 'angle': -45},
|
||||
},
|
||||
}
|
||||
result = _apply_support_filters(data, 'target', support_filters)
|
||||
assert len(result) <= 3
|
||||
assert list(result.columns) == ['a', 'b', 'target']
|
||||
|
||||
def test_missing_upper_or_lower_skips_variable(self):
|
||||
"""If upper_line or lower_line is missing, that variable is skipped."""
|
||||
data = pd.DataFrame({'x': [1, 2, 3], 'target': [10, 20, 30]})
|
||||
support_filters = {
|
||||
'x': {'upper_line': {'intercept': 100, 'angle': 0}},
|
||||
}
|
||||
result = _apply_support_filters(data, 'target', support_filters)
|
||||
pd.testing.assert_frame_equal(result, data)
|
||||
|
||||
Reference in New Issue
Block a user