SIENTIAPDE-1241: refactor train_model workflow due to I/O errors.
This commit is contained in:
@@ -1,497 +0,0 @@
|
||||
"""Unit tests for TrainModelParams class."""
|
||||
|
||||
import pytest
|
||||
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def valid_params_dict():
|
||||
"""Create valid parameters dictionary for testing."""
|
||||
return {
|
||||
'variable_columns': ['var1', 'var2', 'var3'],
|
||||
'lag_train': 5,
|
||||
'lag_val': 3,
|
||||
'target_variable': 'target',
|
||||
'rem_static_win': True,
|
||||
'low_lim': {'var1': 0.0, 'var2': 0.0, 'var3': 0.0},
|
||||
'upp_lim': {'var1': 100.0, 'var2': 100.0, 'var3': 100.0},
|
||||
'window': 10,
|
||||
'use_scaler': True,
|
||||
'include_ar': False,
|
||||
'bucket_name': 'test-bucket',
|
||||
'file_name': 'test-file.csv',
|
||||
'line_separator': '\n',
|
||||
'decimal_separator': '.',
|
||||
'train_size': 80,
|
||||
'shuffle': True,
|
||||
'experiment_run_id': 123,
|
||||
'experiment_name': 'Test experiment name',
|
||||
'removed_intervals': [],
|
||||
}
|
||||
|
||||
|
||||
def test_train_model_params_creation_with_valid_params(valid_params_dict):
|
||||
"""Test creating TrainModelParams with all valid parameters."""
|
||||
params = TrainModelParams(**valid_params_dict)
|
||||
|
||||
assert params.variable_columns == ['var1', 'var2', 'var3']
|
||||
assert params.lag_train == 5
|
||||
assert params.lag_val == 3
|
||||
assert params.target_variable == 'target'
|
||||
assert params.rem_static_win is True
|
||||
assert params.low_lim == {'var1': 0.0, 'var2': 0.0, 'var3': 0.0}
|
||||
assert params.upp_lim == {'var1': 100.0, 'var2': 100.0, 'var3': 100.0}
|
||||
assert params.window == 10
|
||||
assert params.use_scaler is True
|
||||
assert params.include_ar is False
|
||||
assert params.bucket_name == 'test-bucket'
|
||||
assert params.file_name == 'test-file.csv'
|
||||
assert params.line_separator == '\n'
|
||||
assert params.decimal_separator == '.'
|
||||
assert params.train_size == 80
|
||||
assert params.shuffle is True
|
||||
assert params.experiment_run_id == 123
|
||||
assert params.experiment_name == 'Test experiment name'
|
||||
assert params.removed_intervals == []
|
||||
|
||||
|
||||
def test_train_model_params_from_dict_creation(valid_params_dict):
|
||||
"""Test creating TrainModelParams using from_dict method."""
|
||||
params = TrainModelParams.from_dict(valid_params_dict)
|
||||
|
||||
assert params.variable_columns == ['var1', 'var2', 'var3']
|
||||
assert params.lag_train == 5
|
||||
assert params.experiment_run_id == 123
|
||||
|
||||
|
||||
def test_train_model_params_variable_columns_none_raises_error(valid_params_dict):
|
||||
"""Test that None variable_columns raises ValueError."""
|
||||
valid_params_dict['variable_columns'] = None
|
||||
|
||||
with pytest.raises(ValueError, match='variable_columns is required and cannot be None'):
|
||||
TrainModelParams.from_dict(valid_params_dict)
|
||||
|
||||
|
||||
def test_train_model_params_variable_columns_wrong_type_raises_error(valid_params_dict):
|
||||
"""Test that wrong type for variable_columns raises TypeError."""
|
||||
valid_params_dict['variable_columns'] = 'not a list'
|
||||
|
||||
with pytest.raises(TypeError, match='variable_columns must be of type list'):
|
||||
TrainModelParams.from_dict(valid_params_dict)
|
||||
|
||||
|
||||
def test_train_model_params_lag_train_none_raises_error(valid_params_dict):
|
||||
"""Test that None lag_train raises ValueError."""
|
||||
valid_params_dict['lag_train'] = None
|
||||
|
||||
with pytest.raises(ValueError, match='lag_train is required and cannot be None'):
|
||||
TrainModelParams.from_dict(valid_params_dict)
|
||||
|
||||
|
||||
def test_train_model_params_lag_train_wrong_type_raises_error(valid_params_dict):
|
||||
"""Test that wrong type for lag_train raises TypeError."""
|
||||
valid_params_dict['lag_train'] = '5'
|
||||
|
||||
with pytest.raises(TypeError, match='lag_train must be of type int'):
|
||||
TrainModelParams.from_dict(valid_params_dict)
|
||||
|
||||
|
||||
def test_train_model_params_target_variable_none_raises_error(valid_params_dict):
|
||||
"""Test that None target_variable raises ValueError."""
|
||||
valid_params_dict['target_variable'] = None
|
||||
|
||||
with pytest.raises(ValueError, match='target_variable is required and cannot be None'):
|
||||
TrainModelParams.from_dict(valid_params_dict)
|
||||
|
||||
|
||||
def test_train_model_params_target_variable_wrong_type_raises_error(valid_params_dict):
|
||||
"""Test that wrong type for target_variable raises TypeError."""
|
||||
valid_params_dict['target_variable'] = 123
|
||||
|
||||
with pytest.raises(TypeError, match='target_variable must be of type str'):
|
||||
TrainModelParams.from_dict(valid_params_dict)
|
||||
|
||||
|
||||
def test_train_model_params_boolean_fields(valid_params_dict):
|
||||
"""Test boolean fields validation."""
|
||||
# Test rem_static_win
|
||||
valid_params_dict['rem_static_win'] = None
|
||||
with pytest.raises(ValueError, match='rem_static_win is required and cannot be None'):
|
||||
TrainModelParams.from_dict(valid_params_dict)
|
||||
|
||||
valid_params_dict['rem_static_win'] = 'true'
|
||||
with pytest.raises(TypeError, match='rem_static_win must be of type bool'):
|
||||
TrainModelParams.from_dict(valid_params_dict)
|
||||
|
||||
|
||||
def test_train_model_params_dict_fields(valid_params_dict):
|
||||
"""Test dict fields validation."""
|
||||
# Test low_lim
|
||||
valid_params_dict['low_lim'] = None
|
||||
with pytest.raises(ValueError, match='low_lim is required and cannot be None'):
|
||||
TrainModelParams.from_dict(valid_params_dict)
|
||||
|
||||
valid_params_dict['low_lim'] = {'var1': 0.0}
|
||||
valid_params_dict['upp_lim'] = 'not a dict'
|
||||
with pytest.raises(TypeError, match='upp_lim must be of type dict'):
|
||||
TrainModelParams.from_dict(valid_params_dict)
|
||||
|
||||
|
||||
def test_train_model_params_bucket_name_none_raises_error(valid_params_dict):
|
||||
"""Test that None bucket_name raises ValueError."""
|
||||
valid_params_dict['bucket_name'] = None
|
||||
|
||||
with pytest.raises(ValueError, match='bucket_name is required and cannot be None'):
|
||||
TrainModelParams.from_dict(valid_params_dict)
|
||||
|
||||
|
||||
def test_train_model_params_file_name_none_raises_error(valid_params_dict):
|
||||
"""Test that None file_name raises ValueError."""
|
||||
valid_params_dict['file_name'] = None
|
||||
|
||||
with pytest.raises(ValueError, match='file_name is required and cannot be None'):
|
||||
TrainModelParams.from_dict(valid_params_dict)
|
||||
|
||||
|
||||
def test_train_model_params_experiment_run_id_none_raises_error(valid_params_dict):
|
||||
"""Test that None experiment_run_id raises ValueError."""
|
||||
valid_params_dict['experiment_run_id'] = None
|
||||
|
||||
with pytest.raises(ValueError, match='experiment_run_id is required and cannot be None'):
|
||||
TrainModelParams.from_dict(valid_params_dict)
|
||||
|
||||
|
||||
def test_train_model_params_experiment_name_none_raises_error(valid_params_dict):
|
||||
"""Test that None experiment_name raises ValueError."""
|
||||
valid_params_dict['experiment_name'] = None
|
||||
|
||||
with pytest.raises(ValueError, match='experiment_name is required and cannot be None'):
|
||||
TrainModelParams.from_dict(valid_params_dict)
|
||||
|
||||
|
||||
def test_train_model_params_removed_intervals_can_be_none(valid_params_dict):
|
||||
"""Test that removed_intervals can be None (uses _check_type not _check_none)."""
|
||||
valid_params_dict['removed_intervals'] = None
|
||||
|
||||
params = TrainModelParams.from_dict(valid_params_dict)
|
||||
assert params.removed_intervals is None
|
||||
|
||||
|
||||
def test_train_model_params_removed_intervals_wrong_type_raises_error(valid_params_dict):
|
||||
"""Test that wrong type for removed_intervals raises TypeError."""
|
||||
valid_params_dict['removed_intervals'] = 'not a list'
|
||||
|
||||
with pytest.raises(TypeError, match='removed_intervals must be of type list'):
|
||||
TrainModelParams.from_dict(valid_params_dict)
|
||||
|
||||
|
||||
def test_train_model_params_removed_intervals_with_values(valid_params_dict):
|
||||
"""Test removed_intervals with actual interval values."""
|
||||
valid_params_dict['removed_intervals'] = [
|
||||
('2023-01-01', '2023-01-10'),
|
||||
('2023-02-01', '2023-02-05'),
|
||||
]
|
||||
|
||||
params = TrainModelParams.from_dict(valid_params_dict)
|
||||
assert len(params.removed_intervals) == 2
|
||||
assert params.removed_intervals[0] == ('2023-01-01', '2023-01-10')
|
||||
|
||||
|
||||
def test_train_model_params_all_fields_count():
|
||||
"""Test that TrainModelParams has exactly 19 required fields."""
|
||||
import inspect
|
||||
|
||||
sig = inspect.signature(TrainModelParams.__init__)
|
||||
# Subtract 1 for 'self'
|
||||
param_count = len(sig.parameters) - 1
|
||||
assert param_count == 19
|
||||
|
||||
|
||||
def test_train_model_params_with_minimal_valid_data():
|
||||
"""Test creating params with minimal valid data."""
|
||||
params = TrainModelParams(
|
||||
variable_columns=['x'],
|
||||
lag_train=1,
|
||||
lag_val=1,
|
||||
target_variable='y',
|
||||
rem_static_win=False,
|
||||
low_lim={},
|
||||
upp_lim={},
|
||||
window=1,
|
||||
use_scaler=False,
|
||||
include_ar=False,
|
||||
bucket_name='bucket',
|
||||
file_name='file.csv',
|
||||
line_separator='\n',
|
||||
decimal_separator='.',
|
||||
train_size=50,
|
||||
shuffle=False,
|
||||
experiment_run_id=1,
|
||||
experiment_name='name',
|
||||
removed_intervals=[],
|
||||
)
|
||||
|
||||
assert params.variable_columns == ['x']
|
||||
assert params.lag_train == 1
|
||||
assert params.experiment_run_id == 1
|
||||
|
||||
|
||||
def test_train_model_params_check_none_method():
|
||||
"""Test _check_none method behavior."""
|
||||
params_dict = {
|
||||
'variable_columns': ['var1'],
|
||||
'lag_train': 5,
|
||||
'lag_val': 3,
|
||||
'target_variable': 'target',
|
||||
'rem_static_win': True,
|
||||
'low_lim': {},
|
||||
'upp_lim': {},
|
||||
'window': 10,
|
||||
'use_scaler': True,
|
||||
'include_ar': False,
|
||||
'bucket_name': 'bucket',
|
||||
'file_name': 'file.csv',
|
||||
'line_separator': '\n',
|
||||
'decimal_separator': '.',
|
||||
'train_size': 80,
|
||||
'shuffle': True,
|
||||
'experiment_run_id': 123,
|
||||
'experiment_name': 'exp',
|
||||
'removed_intervals': [],
|
||||
}
|
||||
|
||||
params = TrainModelParams(**params_dict)
|
||||
|
||||
# Test that _check_none is a private method
|
||||
assert hasattr(params, '_check_none')
|
||||
assert callable(params._check_none)
|
||||
|
||||
|
||||
def test_train_model_params_check_type_method():
|
||||
"""Test _check_type method behavior."""
|
||||
params_dict = {
|
||||
'variable_columns': ['var1'],
|
||||
'lag_train': 5,
|
||||
'lag_val': 3,
|
||||
'target_variable': 'target',
|
||||
'rem_static_win': True,
|
||||
'low_lim': {},
|
||||
'upp_lim': {},
|
||||
'window': 10,
|
||||
'use_scaler': True,
|
||||
'include_ar': False,
|
||||
'bucket_name': 'bucket',
|
||||
'file_name': 'file.csv',
|
||||
'line_separator': '\n',
|
||||
'decimal_separator': '.',
|
||||
'train_size': 80,
|
||||
'shuffle': True,
|
||||
'experiment_run_id': 123,
|
||||
'experiment_name': 'exp',
|
||||
'removed_intervals': [],
|
||||
}
|
||||
|
||||
params = TrainModelParams(**params_dict)
|
||||
|
||||
# Test that _check_type is a private method
|
||||
assert hasattr(params, '_check_type')
|
||||
assert callable(params._check_type)
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Tests for validate_business_rules method
|
||||
# ============================================================================
|
||||
|
||||
|
||||
def test_validate_business_rules_success(valid_params_dict):
|
||||
"""Test that valid params pass business rules validation."""
|
||||
# Ensure target_variable is in variable_columns
|
||||
valid_params_dict['target_variable'] = 'var1'
|
||||
|
||||
params = TrainModelParams.from_dict(valid_params_dict)
|
||||
|
||||
# Should not raise any exception
|
||||
params.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_train_size_too_low(valid_params_dict):
|
||||
"""Test that train_size < 1 raises ValueError."""
|
||||
valid_params_dict['train_size'] = 0
|
||||
valid_params_dict['target_variable'] = 'var1'
|
||||
|
||||
params = TrainModelParams.from_dict(valid_params_dict)
|
||||
|
||||
with pytest.raises(ValueError, match='train_size must be between 1 and 99'):
|
||||
params.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_train_size_too_high(valid_params_dict):
|
||||
"""Test that train_size > 99 raises ValueError."""
|
||||
valid_params_dict['train_size'] = 100
|
||||
valid_params_dict['target_variable'] = 'var1'
|
||||
|
||||
params = TrainModelParams.from_dict(valid_params_dict)
|
||||
|
||||
with pytest.raises(ValueError, match='train_size must be between 1 and 99'):
|
||||
params.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_empty_variable_columns(valid_params_dict):
|
||||
"""Test that empty variable_columns raises ValueError."""
|
||||
valid_params_dict['variable_columns'] = []
|
||||
|
||||
params = TrainModelParams.from_dict(valid_params_dict)
|
||||
|
||||
with pytest.raises(ValueError, match='variable_columns cannot be empty'):
|
||||
params.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_lag_train_zero(valid_params_dict):
|
||||
"""Test that lag_train = 0 raises ValueError."""
|
||||
valid_params_dict['lag_train'] = 0
|
||||
valid_params_dict['target_variable'] = 'var1'
|
||||
|
||||
params = TrainModelParams.from_dict(valid_params_dict)
|
||||
|
||||
with pytest.raises(ValueError, match='lag_train must be positive'):
|
||||
params.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_lag_train_negative(valid_params_dict):
|
||||
"""Test that lag_train < 0 raises ValueError."""
|
||||
valid_params_dict['lag_train'] = -1
|
||||
valid_params_dict['target_variable'] = 'var1'
|
||||
|
||||
params = TrainModelParams.from_dict(valid_params_dict)
|
||||
|
||||
with pytest.raises(ValueError, match='lag_train must be positive'):
|
||||
params.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_lag_val_zero(valid_params_dict):
|
||||
"""Test that lag_val = 0 raises ValueError."""
|
||||
valid_params_dict['lag_val'] = 0
|
||||
valid_params_dict['target_variable'] = 'var1'
|
||||
|
||||
params = TrainModelParams.from_dict(valid_params_dict)
|
||||
|
||||
with pytest.raises(ValueError, match='lag_val must be positive'):
|
||||
params.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_window_zero(valid_params_dict):
|
||||
"""Test that window = 0 raises ValueError."""
|
||||
valid_params_dict['window'] = 0
|
||||
valid_params_dict['target_variable'] = 'var1'
|
||||
|
||||
params = TrainModelParams.from_dict(valid_params_dict)
|
||||
|
||||
with pytest.raises(ValueError, match='window must be positive'):
|
||||
params.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_low_lim_upp_lim_keys_mismatch(valid_params_dict):
|
||||
"""Test that mismatched keys in low_lim and upp_lim raises ValueError."""
|
||||
valid_params_dict['low_lim'] = {'var1': 0.0, 'var2': 0.0}
|
||||
valid_params_dict['upp_lim'] = {'var1': 100.0, 'var3': 100.0} # var3 instead of var2
|
||||
valid_params_dict['target_variable'] = 'var1'
|
||||
|
||||
params = TrainModelParams.from_dict(valid_params_dict)
|
||||
|
||||
with pytest.raises(ValueError, match='low_lim and upp_lim must have the same keys'):
|
||||
params.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_low_lim_greater_than_upp_lim(valid_params_dict):
|
||||
"""Test that low_lim >= upp_lim raises ValueError."""
|
||||
valid_params_dict['low_lim'] = {'var1': 100.0, 'var2': 0.0, 'var3': 0.0}
|
||||
valid_params_dict['upp_lim'] = {'var1': 50.0, 'var2': 100.0, 'var3': 100.0}
|
||||
valid_params_dict['target_variable'] = 'var2'
|
||||
|
||||
params = TrainModelParams.from_dict(valid_params_dict)
|
||||
|
||||
with pytest.raises(ValueError, match='low_lim must be less than upp_lim for variable "var1"'):
|
||||
params.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_low_lim_equal_to_upp_lim(valid_params_dict):
|
||||
"""Test that low_lim == upp_lim raises ValueError."""
|
||||
valid_params_dict['low_lim'] = {'var1': 50.0, 'var2': 0.0, 'var3': 0.0}
|
||||
valid_params_dict['upp_lim'] = {'var1': 50.0, 'var2': 100.0, 'var3': 100.0}
|
||||
valid_params_dict['target_variable'] = 'var2'
|
||||
|
||||
params = TrainModelParams.from_dict(valid_params_dict)
|
||||
|
||||
with pytest.raises(ValueError, match='low_lim must be less than upp_lim for variable "var1"'):
|
||||
params.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_target_not_in_variable_columns(valid_params_dict):
|
||||
"""Test that target_variable not in variable_columns raises ValueError."""
|
||||
valid_params_dict['target_variable'] = 'nonexistent_var'
|
||||
|
||||
params = TrainModelParams.from_dict(valid_params_dict)
|
||||
|
||||
with pytest.raises(
|
||||
ValueError, match='target_variable "nonexistent_var" must be in variable_columns'
|
||||
):
|
||||
params.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_empty_bucket_name(valid_params_dict):
|
||||
"""Test that empty bucket_name raises ValueError."""
|
||||
valid_params_dict['bucket_name'] = ' '
|
||||
valid_params_dict['target_variable'] = 'var1'
|
||||
|
||||
params = TrainModelParams.from_dict(valid_params_dict)
|
||||
|
||||
with pytest.raises(ValueError, match='bucket_name cannot be empty or whitespace'):
|
||||
params.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_empty_file_name(valid_params_dict):
|
||||
"""Test that empty file_name raises ValueError."""
|
||||
valid_params_dict['file_name'] = ''
|
||||
valid_params_dict['target_variable'] = 'var1'
|
||||
|
||||
params = TrainModelParams.from_dict(valid_params_dict)
|
||||
|
||||
with pytest.raises(ValueError, match='file_name cannot be empty or whitespace'):
|
||||
params.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_empty_experiment_name(valid_params_dict):
|
||||
"""Test that empty experiment_name raises ValueError."""
|
||||
valid_params_dict['experiment_name'] = ' '
|
||||
valid_params_dict['target_variable'] = 'var1'
|
||||
|
||||
params = TrainModelParams.from_dict(valid_params_dict)
|
||||
|
||||
with pytest.raises(ValueError, match='experiment_name cannot be empty or whitespace'):
|
||||
params.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_all_valid_edge_cases(valid_params_dict):
|
||||
"""Test that edge case valid values pass validation."""
|
||||
valid_params_dict['train_size'] = 1 # Minimum valid
|
||||
valid_params_dict['lag_train'] = 1 # Minimum valid
|
||||
valid_params_dict['lag_val'] = 1 # Minimum valid
|
||||
valid_params_dict['window'] = 1 # Minimum valid
|
||||
valid_params_dict['target_variable'] = 'var1'
|
||||
|
||||
params = TrainModelParams.from_dict(valid_params_dict)
|
||||
|
||||
# Should not raise any exception
|
||||
params.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_train_size_99(valid_params_dict):
|
||||
"""Test that train_size = 99 (maximum valid) passes validation."""
|
||||
valid_params_dict['train_size'] = 99
|
||||
valid_params_dict['target_variable'] = 'var1'
|
||||
|
||||
params = TrainModelParams.from_dict(valid_params_dict)
|
||||
|
||||
# Should not raise any exception
|
||||
params.validate_business_rules()
|
||||
Reference in New Issue
Block a user