SIENTIAPDE-1430: Introduce static_threshold parameter for static window removal.
This parameter allows customizing the threshold (1-1000) used when rem_static_win is enabled, defaulting to 1 if null. Updates include parameter definition, business rule validation, repository logic for passing the threshold, documentation in README.md and PIPELINE_PARAMS_CHANGELOG.md, and new unit and integration tests.
This commit is contained in:
@@ -34,6 +34,7 @@ def valid_train_params_dict():
|
||||
'end_date': None,
|
||||
'scaler_name': 'Standard Scaler',
|
||||
'support_filters': {},
|
||||
'static_threshold': None,
|
||||
}
|
||||
|
||||
|
||||
@@ -564,3 +565,96 @@ def test_validate_business_rules_polynomial_regression_valid(valid_train_params_
|
||||
params = TrainModelParams.from_dict(valid_train_params_dict)
|
||||
|
||||
params.validate_business_rules() # Should not raise
|
||||
|
||||
|
||||
def test_validate_business_rules_static_threshold_valid(valid_train_params_dict):
|
||||
"""Test validate_business_rules accepts valid static_threshold when rem_static_win is True."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
valid_train_params_dict['rem_static_win'] = True
|
||||
valid_train_params_dict['static_threshold'] = 500
|
||||
params = TrainModelParams.from_dict(valid_train_params_dict)
|
||||
|
||||
params.validate_business_rules() # Should not raise
|
||||
|
||||
|
||||
def test_validate_business_rules_static_threshold_min_valid(valid_train_params_dict):
|
||||
"""Test validate_business_rules accepts static_threshold = 1."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
valid_train_params_dict['rem_static_win'] = True
|
||||
valid_train_params_dict['static_threshold'] = 1
|
||||
params = TrainModelParams.from_dict(valid_train_params_dict)
|
||||
|
||||
params.validate_business_rules() # Should not raise
|
||||
|
||||
|
||||
def test_validate_business_rules_static_threshold_max_valid(valid_train_params_dict):
|
||||
"""Test validate_business_rules accepts static_threshold = 1000."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
valid_train_params_dict['rem_static_win'] = True
|
||||
valid_train_params_dict['static_threshold'] = 1000
|
||||
params = TrainModelParams.from_dict(valid_train_params_dict)
|
||||
|
||||
params.validate_business_rules() # Should not raise
|
||||
|
||||
|
||||
def test_validate_business_rules_static_threshold_below_min(valid_train_params_dict):
|
||||
"""Test validate_business_rules raises error when static_threshold < 1."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
valid_train_params_dict['rem_static_win'] = True
|
||||
valid_train_params_dict['static_threshold'] = 0
|
||||
params = TrainModelParams.from_dict(valid_train_params_dict)
|
||||
|
||||
with pytest.raises(ValueError, match='static_threshold must be between 1 and 1000, got 0'):
|
||||
params.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_static_threshold_above_max(valid_train_params_dict):
|
||||
"""Test validate_business_rules raises error when static_threshold > 1000."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
valid_train_params_dict['rem_static_win'] = True
|
||||
valid_train_params_dict['static_threshold'] = 1001
|
||||
params = TrainModelParams.from_dict(valid_train_params_dict)
|
||||
|
||||
with pytest.raises(ValueError, match='static_threshold must be between 1 and 1000, got 1001'):
|
||||
params.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_static_threshold_none_when_rem_static_win_true(
|
||||
valid_train_params_dict,
|
||||
):
|
||||
"""Test validate_business_rules accepts None static_threshold when rem_static_win is True."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
valid_train_params_dict['rem_static_win'] = True
|
||||
valid_train_params_dict['static_threshold'] = None
|
||||
params = TrainModelParams.from_dict(valid_train_params_dict)
|
||||
|
||||
params.validate_business_rules() # Should not raise - None is allowed
|
||||
|
||||
|
||||
def test_validate_business_rules_static_threshold_ignored_when_rem_static_win_false(
|
||||
valid_train_params_dict,
|
||||
):
|
||||
"""Test validate_business_rules ignores static_threshold when rem_static_win is False."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
valid_train_params_dict['rem_static_win'] = False
|
||||
valid_train_params_dict['static_threshold'] = 5000 # Invalid value, but should be ignored
|
||||
params = TrainModelParams.from_dict(valid_train_params_dict)
|
||||
|
||||
params.validate_business_rules() # Should not raise - validation skipped
|
||||
|
||||
|
||||
def test_from_dict_static_threshold_type_error(valid_train_params_dict):
|
||||
"""Test from_dict raises TypeError when static_threshold has wrong type."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
valid_train_params_dict['static_threshold'] = 'not_an_int'
|
||||
|
||||
with pytest.raises(TypeError, match='static_threshold must be of type int, but got str'):
|
||||
TrainModelParams.from_dict(valid_train_params_dict)
|
||||
|
||||
@@ -40,6 +40,7 @@ def sample_params():
|
||||
end_date=None,
|
||||
scaler_name='Standard Scaler',
|
||||
support_filters={},
|
||||
static_threshold=None,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -55,6 +55,8 @@ def mock_train_result():
|
||||
result.params.include_ar = False
|
||||
result.params.train_size = 80
|
||||
result.params.removed_intervals = []
|
||||
result.params.rem_static_win = True
|
||||
result.params.static_threshold = None
|
||||
result.run_name = 'test_run'
|
||||
result.run_dir = '/tmp/test_run' # noqa: S108
|
||||
result.report_path = '/tmp/test_run/report.html' # noqa: S108
|
||||
|
||||
@@ -56,6 +56,7 @@ def sample_params():
|
||||
end_date=None,
|
||||
scaler_name='None',
|
||||
support_filters={},
|
||||
static_threshold=None,
|
||||
)
|
||||
|
||||
|
||||
@@ -137,6 +138,7 @@ class TestExtractModelEquation:
|
||||
end_date=None,
|
||||
scaler_name='None',
|
||||
support_filters={},
|
||||
static_threshold=None,
|
||||
)
|
||||
|
||||
# Mock model with single coefficient
|
||||
@@ -221,12 +223,33 @@ class TestInitDataPreprocessor:
|
||||
assert preprocessor.ar_var is None
|
||||
|
||||
def test_init_preprocessor_with_static_removal(self, training_repo, sample_params):
|
||||
"""Test preprocessor with static window removal enabled."""
|
||||
"""Test preprocessor with static window removal enabled and no static_threshold."""
|
||||
sample_params.rem_static_win = True
|
||||
sample_params.static_threshold = None
|
||||
preprocessor = training_repo._init_data_preprocessor(sample_params)
|
||||
|
||||
assert preprocessor.static_threshold == 1
|
||||
|
||||
def test_init_preprocessor_with_static_removal_custom_threshold(
|
||||
self, training_repo, sample_params
|
||||
):
|
||||
"""Test preprocessor with static window removal and custom static_threshold."""
|
||||
sample_params.rem_static_win = True
|
||||
sample_params.static_threshold = 500
|
||||
preprocessor = training_repo._init_data_preprocessor(sample_params)
|
||||
|
||||
assert preprocessor.static_threshold == 500
|
||||
|
||||
def test_init_preprocessor_without_static_removal_ignores_threshold(
|
||||
self, training_repo, sample_params
|
||||
):
|
||||
"""Test preprocessor without static removal ignores static_threshold."""
|
||||
sample_params.rem_static_win = False
|
||||
sample_params.static_threshold = 500
|
||||
preprocessor = training_repo._init_data_preprocessor(sample_params)
|
||||
|
||||
assert preprocessor.static_threshold is None
|
||||
|
||||
def test_init_preprocessor_lag_configuration(self, training_repo, sample_params):
|
||||
"""Test preprocessor lag configuration."""
|
||||
sample_params.lag_train = {'var1': 5, 'var2': 5, 'var3': 5}
|
||||
|
||||
Reference in New Issue
Block a user