SIENTIAPDE-1579: Updated tests

This commit is contained in:
Kou Kinoshita
2026-02-18 15:05:37 -03:00
parent b6eb346a0b
commit c6d6f94e05
5 changed files with 179 additions and 33 deletions

View File

@@ -6,12 +6,13 @@ Steps performed (default):
2. Insert a new experiment_run record in Postgres and capture the generated ID. 2. Insert a new experiment_run record in Postgres and capture the generated ID.
3. Trigger the Temporal `train_model` workflow with the correct payload. 3. Trigger the Temporal `train_model` workflow with the correct payload.
Local diagnosis (--local / --validate-only): Alternatives for local diagnosis (--local / --validate-only):
- example: python scripts/run_training_test.py --scenario 01-linear-regression-basic --local --csv docs/test-model-data.csv
- --validate-only: Validates scenario parameters only (no MinIO, Postgres, Temporal). - --validate-only: Validates scenario parameters only (no MinIO, Postgres, Temporal).
- --local: Runs the same training pipeline locally (validate + load CSV + train + - --local: Runs the same training pipeline locally (validate + load CSV + train +
after_train_calculation). Use to get full Python tracebacks for debugging. after_train_calculation). Use to get full Python tracebacks for debugging.
Does not upload to MinIO, insert DB, or start Temporal. By default skips Does not upload to MinIO, insert DB, or start Temporal.
MLflow save; use --local-save-mlflow to also test saving to MLflow. By default skips MLflow save; use --local-save-mlflow to also test saving to MLflow.
Prerequisites (default flow): Prerequisites (default flow):
- `mc` CLI configured with alias defined in MINIO_ALIAS. - `mc` CLI configured with alias defined in MINIO_ALIAS.

View File

@@ -41,6 +41,8 @@ def sample_params():
scaler_name='Standard Scaler', scaler_name='Standard Scaler',
support_filters={}, support_filters={},
static_threshold=None, static_threshold=None,
date_column=None,
date_format=None,
) )

View File

@@ -344,9 +344,14 @@ def test_create_run_directory_success(
result = repo._create_run_directory('/tmp/reports', 'test_run') # noqa: S108 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 expected_path = os.path.normpath(
assert result == expected_path os.path.join('/tmp/reports', 'temp', 'test_run_20240101_120000_123456') # noqa: S108
mock_makedirs.assert_called_once_with(expected_path, exist_ok=True) )
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') @patch('model_manager.utils.repository.model_repository.ModelServing')

View File

@@ -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)