|
|
|
@@ -10,7 +10,10 @@ import pytest
|
|
|
|
from model_manager.sientia.models import LinearRegressionModel
|
|
|
|
from model_manager.sientia.models import LinearRegressionModel
|
|
|
|
from model_manager.utils.models.train_model_params import TrainModelParams
|
|
|
|
from model_manager.utils.models.train_model_params import TrainModelParams
|
|
|
|
from model_manager.utils.models.train_model_result import TrainModelResult
|
|
|
|
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
|
|
|
|
@pytest.fixture
|
|
|
|
@@ -57,6 +60,8 @@ def sample_params():
|
|
|
|
scaler_name='None',
|
|
|
|
scaler_name='None',
|
|
|
|
support_filters={},
|
|
|
|
support_filters={},
|
|
|
|
static_threshold=None,
|
|
|
|
static_threshold=None,
|
|
|
|
|
|
|
|
date_column=None,
|
|
|
|
|
|
|
|
date_format=None,
|
|
|
|
)
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
@@ -139,6 +144,8 @@ class TestExtractModelEquation:
|
|
|
|
scaler_name='None',
|
|
|
|
scaler_name='None',
|
|
|
|
support_filters={},
|
|
|
|
support_filters={},
|
|
|
|
static_threshold=None,
|
|
|
|
static_threshold=None,
|
|
|
|
|
|
|
|
date_column=None,
|
|
|
|
|
|
|
|
date_format=None,
|
|
|
|
)
|
|
|
|
)
|
|
|
|
|
|
|
|
|
|
|
|
# Mock model with single coefficient
|
|
|
|
# Mock model with single coefficient
|
|
|
|
@@ -585,7 +592,7 @@ class TestTrain:
|
|
|
|
"""
|
|
|
|
"""
|
|
|
|
return BytesIO(csv_content.encode('utf-8'))
|
|
|
|
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.split_train_test')
|
|
|
|
@patch('model_manager.utils.repository.training_repository.load_data')
|
|
|
|
@patch('model_manager.utils.repository.training_repository.load_data')
|
|
|
|
def test_train_basic_workflow(
|
|
|
|
def test_train_basic_workflow(
|
|
|
|
@@ -631,7 +638,7 @@ class TestTrain:
|
|
|
|
# Verify split was called
|
|
|
|
# Verify split was called
|
|
|
|
assert mock_split_train_test.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.split_train_test')
|
|
|
|
@patch('model_manager.utils.repository.training_repository.load_data')
|
|
|
|
@patch('model_manager.utils.repository.training_repository.load_data')
|
|
|
|
def test_train_with_scaler(
|
|
|
|
def test_train_with_scaler(
|
|
|
|
@@ -661,7 +668,7 @@ class TestTrain:
|
|
|
|
assert result is not None
|
|
|
|
assert result is not None
|
|
|
|
assert result.scaler_dict 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.split_train_test')
|
|
|
|
@patch('model_manager.utils.repository.training_repository.load_data')
|
|
|
|
@patch('model_manager.utils.repository.training_repository.load_data')
|
|
|
|
def test_train_with_shuffle_enabled(
|
|
|
|
def test_train_with_shuffle_enabled(
|
|
|
|
@@ -692,7 +699,7 @@ class TestTrain:
|
|
|
|
call_kwargs = mock_split_train_test.call_args[1]
|
|
|
|
call_kwargs = mock_split_train_test.call_args[1]
|
|
|
|
assert call_kwargs['shuffle'] is True
|
|
|
|
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.split_train_test')
|
|
|
|
@patch('model_manager.utils.repository.training_repository.load_data')
|
|
|
|
@patch('model_manager.utils.repository.training_repository.load_data')
|
|
|
|
def test_train_with_different_train_size(
|
|
|
|
def test_train_with_different_train_size(
|
|
|
|
@@ -723,7 +730,7 @@ class TestTrain:
|
|
|
|
call_kwargs = mock_split_train_test.call_args[1]
|
|
|
|
call_kwargs = mock_split_train_test.call_args[1]
|
|
|
|
assert call_kwargs['train_size'] == 0.7
|
|
|
|
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.split_train_test')
|
|
|
|
@patch('model_manager.utils.repository.training_repository.load_data')
|
|
|
|
@patch('model_manager.utils.repository.training_repository.load_data')
|
|
|
|
def test_train_raises_on_empty_data_after_transform(
|
|
|
|
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'):
|
|
|
|
with pytest.raises(ValueError, match='Data view is empty after transformation'):
|
|
|
|
training_repo.train(sample_csv_data, sample_params)
|
|
|
|
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.split_train_test')
|
|
|
|
@patch('model_manager.utils.repository.training_repository.load_data')
|
|
|
|
@patch('model_manager.utils.repository.training_repository.load_data')
|
|
|
|
def test_train_logs_success(
|
|
|
|
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
|
|
|
|
'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.split_train_test')
|
|
|
|
@patch('model_manager.utils.repository.training_repository.load_data')
|
|
|
|
@patch('model_manager.utils.repository.training_repository.load_data')
|
|
|
|
def test_train_with_custom_separators(
|
|
|
|
def test_train_with_custom_separators(
|
|
|
|
@@ -812,7 +819,7 @@ class TestTrain:
|
|
|
|
# Verify load_data was called with custom separators
|
|
|
|
# Verify load_data was called with custom separators
|
|
|
|
mock_load_data.assert_called_once_with(sample_csv_data, ';', ',')
|
|
|
|
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.split_train_test')
|
|
|
|
@patch('model_manager.utils.repository.training_repository.load_data')
|
|
|
|
@patch('model_manager.utils.repository.training_repository.load_data')
|
|
|
|
def test_train_result_contains_all_fields(
|
|
|
|
def test_train_result_contains_all_fields(
|
|
|
|
@@ -862,19 +869,55 @@ class TestConfigureDatetimeIndex:
|
|
|
|
"""Create a TrainingRepository instance."""
|
|
|
|
"""Create a TrainingRepository instance."""
|
|
|
|
return TrainingRepository(logger=mock_logger)
|
|
|
|
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."""
|
|
|
|
"""Test _configure_datetime_index when index is already DatetimeIndex."""
|
|
|
|
data = pd.DataFrame(
|
|
|
|
data = pd.DataFrame(
|
|
|
|
{'var1': [1, 2, 3], 'var2': [4, 5, 6]},
|
|
|
|
{'var1': [1, 2, 3], 'var2': [4, 5, 6]},
|
|
|
|
index=pd.to_datetime(['2023-01-01', '2023-01-02', '2023-01-03']),
|
|
|
|
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 isinstance(result.index, pd.DatetimeIndex)
|
|
|
|
assert len(result) == 3
|
|
|
|
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."""
|
|
|
|
"""Test _configure_datetime_index with 'timestamp' column."""
|
|
|
|
data = pd.DataFrame(
|
|
|
|
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 isinstance(result.index, pd.DatetimeIndex)
|
|
|
|
assert 'timestamp' not in result.columns
|
|
|
|
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."""
|
|
|
|
"""Test _configure_datetime_index with 'date' column."""
|
|
|
|
data = pd.DataFrame(
|
|
|
|
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 isinstance(result.index, pd.DatetimeIndex)
|
|
|
|
assert 'date' not in result.columns
|
|
|
|
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."""
|
|
|
|
"""Test _configure_datetime_index with 'datetime' column."""
|
|
|
|
data = pd.DataFrame(
|
|
|
|
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 isinstance(result.index, pd.DatetimeIndex)
|
|
|
|
assert 'datetime' not in result.columns
|
|
|
|
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."""
|
|
|
|
"""Test _configure_datetime_index when first column looks like datetime."""
|
|
|
|
data = pd.DataFrame(
|
|
|
|
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 isinstance(result.index, pd.DatetimeIndex)
|
|
|
|
assert 'my_date' not in result.columns
|
|
|
|
assert 'my_date' not in result.columns
|
|
|
|
|
|
|
|
|
|
|
|
@pytest.mark.filterwarnings('ignore::UserWarning')
|
|
|
|
@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."""
|
|
|
|
"""Test _configure_datetime_index when no timestamp column found."""
|
|
|
|
data = pd.DataFrame(
|
|
|
|
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)
|
|
|
|
# Should return original data unchanged (no valid datetime columns)
|
|
|
|
assert 'var1' in result.columns
|
|
|
|
assert 'var1' in result.columns
|
|
|
|
assert 'var2' in result.columns
|
|
|
|
assert 'var2' in result.columns
|
|
|
|
|
|
|
|
|
|
|
|
@pytest.mark.filterwarnings('ignore::UserWarning')
|
|
|
|
@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."""
|
|
|
|
"""Test _configure_datetime_index with invalid timestamp values."""
|
|
|
|
data = pd.DataFrame(
|
|
|
|
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
|
|
|
|
# Should skip invalid column and try first column
|
|
|
|
assert 'var1' in result.columns
|
|
|
|
assert 'var1' in result.columns
|
|
|
|
|
|
|
|
|
|
|
|
@pytest.mark.filterwarnings('ignore::UserWarning')
|
|
|
|
@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."""
|
|
|
|
"""Test _configure_datetime_index when first column is not datetime."""
|
|
|
|
data = pd.DataFrame(
|
|
|
|
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
|
|
|
|
# Should return original data unchanged
|
|
|
|
assert 'var1' in result.columns
|
|
|
|
assert 'var1' in result.columns
|
|
|
|
assert 'var2' 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."""
|
|
|
|
"""Test _configure_datetime_index when first column has all NaN values."""
|
|
|
|
data = pd.DataFrame(
|
|
|
|
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)
|
|
|
|
# Should return original data unchanged (first column has no valid values)
|
|
|
|
assert 'first_col' in result.columns
|
|
|
|
assert 'first_col' in result.columns
|
|
|
|
@@ -1072,3 +1115,98 @@ class TestExtractModelEquationPolynomial:
|
|
|
|
|
|
|
|
|
|
|
|
assert 'equation_string' in result
|
|
|
|
assert 'equation_string' in result
|
|
|
|
assert 'latex_equation' 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)
|
|
|
|
|