feat: enhance training and experiment tracking functionality
- Updated `Activities` class to improve garbage collection handling. - Enhanced error messaging in `ExperimentTracking` for better clarity on update failures. - Refactored `Training` class to streamline exception handling and improve type hints. - Introduced new methods in `TrainModelParams` for better handling of experiment run IDs and model metadata. - Added functionality to extract model equations in `DataManagerRepository` for linear regression models.
This commit is contained in:
@@ -1,668 +1,262 @@
|
||||
"""Unit tests for TrainModelParams with 100% coverage."""
|
||||
"""Unit tests for TrainModelParams (current schema)."""
|
||||
|
||||
import copy
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def valid_train_params_dict():
|
||||
"""Create a valid dictionary for TrainModelParams."""
|
||||
def minimal_model_metadata() -> dict:
|
||||
"""Minimal truthy metadata so validate_business_rules passes schema lookup."""
|
||||
return {'schemas': {'components': {'schemas': {}}}}
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def valid_train_params_dict(minimal_model_metadata) -> dict:
|
||||
"""Valid dictionary for TrainModelParams.from_dict."""
|
||||
return {
|
||||
'variable_columns': ['var1', 'var2'],
|
||||
'lag_train': {'var1': 5, 'var2': 5},
|
||||
'lag_val': {'var1': 3, 'var2': 3},
|
||||
'target_variable': 'target',
|
||||
'rem_static_win': True,
|
||||
'low_lim': {'var1': 0.0, 'var2': 1.0},
|
||||
'upp_lim': {'var1': 10.0, 'var2': 20.0},
|
||||
'window': 10,
|
||||
'use_scaler': True,
|
||||
'include_ar': False,
|
||||
'bucket_name': 'test-bucket',
|
||||
'file_name': 'test-file.csv',
|
||||
'line_separator': ',',
|
||||
'decimal_separator': '.',
|
||||
'date_column': None,
|
||||
'date_format': None,
|
||||
'train_size': 80,
|
||||
'shuffle': True,
|
||||
'random_state': 42,
|
||||
'experiment_run_id': 1,
|
||||
'removed_intervals': [],
|
||||
'model_name': 'Linear Regression',
|
||||
'degree': 1,
|
||||
'interaction_only': False,
|
||||
'nan_treatment': 'drop',
|
||||
'start_date': None,
|
||||
'end_date': None,
|
||||
'scaler_name': 'Standard Scaler',
|
||||
'support_filters': {},
|
||||
'static_threshold': None,
|
||||
'val_file_name': None,
|
||||
'data_model_kwargs': {},
|
||||
'model_kwargs': {},
|
||||
'opt_params': {},
|
||||
'model_type': 'linear_regression',
|
||||
'model_id': None,
|
||||
'model_metadata': minimal_model_metadata,
|
||||
}
|
||||
|
||||
|
||||
def test_train_model_params_from_dict_success(valid_train_params_dict):
|
||||
"""Test TrainModelParams.from_dict with valid data."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
def test_from_dict_success(valid_train_params_dict):
|
||||
"""from_dict builds params and experiment_name from model_name."""
|
||||
params = TrainModelParams.from_dict(valid_train_params_dict)
|
||||
|
||||
assert params.variable_columns == ['var1', 'var2']
|
||||
assert params.lag_train == {'var1': 5, 'var2': 5}
|
||||
assert params.lag_val == {'var1': 3, 'var2': 3}
|
||||
assert params.target_variable == 'target'
|
||||
assert params.rem_static_win is True
|
||||
assert params.low_lim == {'var1': 0.0, 'var2': 1.0}
|
||||
assert params.upp_lim == {'var1': 10.0, 'var2': 20.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 == ','
|
||||
assert params.decimal_separator == '.'
|
||||
assert params.train_size == 80
|
||||
assert params.shuffle is True
|
||||
assert params.experiment_run_id == 1
|
||||
assert params.removed_intervals == []
|
||||
assert params.experiment_name == 'Linear Regression_experiment'
|
||||
assert params.model_metadata is valid_train_params_dict['model_metadata']
|
||||
|
||||
|
||||
def test_train_model_params_check_none_raises_value_error():
|
||||
"""Test _check_none raises ValueError when value is None."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
def test_from_dict_coerces_experiment_run_id_string(valid_train_params_dict):
|
||||
"""Numeric string experiment_run_id is coerced to int."""
|
||||
d = copy.deepcopy(valid_train_params_dict)
|
||||
d['experiment_run_id'] = '42'
|
||||
params = TrainModelParams.from_dict(d)
|
||||
assert params.experiment_run_id == 42
|
||||
|
||||
with pytest.raises(ValueError, match='test_field is required and cannot be None'):
|
||||
|
||||
def test_from_dict_model_metadata_none(valid_train_params_dict):
|
||||
"""model_metadata may be None before load_model_metadata activity."""
|
||||
d = copy.deepcopy(valid_train_params_dict)
|
||||
d['model_metadata'] = None
|
||||
params = TrainModelParams.from_dict(d)
|
||||
assert params.model_metadata is None
|
||||
|
||||
|
||||
def test_coerce_experiment_run_id_rejects_bool():
|
||||
"""Boolean must not be accepted as experiment_run_id."""
|
||||
with pytest.raises(TypeError, match='experiment_run_id must be an integer'):
|
||||
TrainModelParams._coerce_experiment_run_id(True)
|
||||
|
||||
|
||||
def test_parse_optional_model_metadata_rejects_list():
|
||||
"""model_metadata must be dict or None."""
|
||||
with pytest.raises(TypeError, match='model_metadata must be a dict or None'):
|
||||
TrainModelParams._parse_optional_model_metadata([])
|
||||
|
||||
|
||||
def test_check_none_raises_value_error():
|
||||
with pytest.raises(ValueError, match='test_field is required'):
|
||||
TrainModelParams._check_none(None, str, 'test_field')
|
||||
|
||||
|
||||
def test_train_model_params_check_none_raises_type_error():
|
||||
"""Test _check_none raises TypeError when type is incorrect."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
with pytest.raises(TypeError, match='test_field must be of type str, but got int'):
|
||||
def test_check_none_raises_type_error():
|
||||
with pytest.raises(TypeError, match='test_field must be of type str'):
|
||||
TrainModelParams._check_none(123, str, 'test_field')
|
||||
|
||||
|
||||
def test_train_model_params_check_none_success():
|
||||
"""Test _check_none returns value when valid."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
result = TrainModelParams._check_none('test_value', str, 'test_field')
|
||||
assert result == 'test_value'
|
||||
|
||||
|
||||
def test_train_model_params_check_type_raises_type_error():
|
||||
"""Test _check_type raises TypeError when type is incorrect."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
with pytest.raises(TypeError, match='test_field must be of type int, but got str'):
|
||||
TrainModelParams._check_type('not_an_int', int, 'test_field')
|
||||
|
||||
|
||||
def test_train_model_params_check_type_success():
|
||||
"""Test _check_type returns value when valid."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
result = TrainModelParams._check_type(42, int, 'test_field')
|
||||
assert result == 42
|
||||
|
||||
|
||||
def test_train_model_params_check_type_with_none():
|
||||
"""Test _check_type allows None value."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
result = TrainModelParams._check_type(None, str, 'test_field')
|
||||
assert result is None
|
||||
|
||||
|
||||
def test_train_model_params_from_dict_missing_field(valid_train_params_dict):
|
||||
"""Test from_dict raises ValueError when required field is missing."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
del valid_train_params_dict['variable_columns']
|
||||
|
||||
with pytest.raises(ValueError, match='variable_columns is required and cannot be None'):
|
||||
TrainModelParams.from_dict(valid_train_params_dict)
|
||||
|
||||
|
||||
def test_train_model_params_from_dict_wrong_type(valid_train_params_dict):
|
||||
"""Test from_dict raises TypeError when field has wrong type."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
valid_train_params_dict['lag_train'] = 'not_a_dict'
|
||||
|
||||
with pytest.raises(TypeError, match='lag_train must be of type dict, but got str'):
|
||||
TrainModelParams.from_dict(valid_train_params_dict)
|
||||
|
||||
|
||||
def test_train_model_params_from_dict_with_none_removed_intervals(valid_train_params_dict):
|
||||
"""Test from_dict allows None for removed_intervals."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
valid_train_params_dict['removed_intervals'] = None
|
||||
|
||||
params = TrainModelParams.from_dict(valid_train_params_dict)
|
||||
assert params.removed_intervals is None
|
||||
|
||||
|
||||
def test_validate_business_rules_success(valid_train_params_dict):
|
||||
"""Test validate_business_rules with valid parameters."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
params = TrainModelParams.from_dict(valid_train_params_dict)
|
||||
params.validate_business_rules() # Should not raise
|
||||
params.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_train_size_too_low(valid_train_params_dict):
|
||||
"""Test validate_business_rules raises error when train_size < 10."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
valid_train_params_dict['train_size'] = 5
|
||||
params = TrainModelParams.from_dict(valid_train_params_dict)
|
||||
|
||||
with pytest.raises(ValueError, match='train_size must be between 10 and 100, got 5'):
|
||||
def test_validate_business_rules_missing_model_metadata(valid_train_params_dict):
|
||||
d = copy.deepcopy(valid_train_params_dict)
|
||||
d['model_metadata'] = None
|
||||
params = TrainModelParams.from_dict(d)
|
||||
with pytest.raises(ValueError, match='model_metadata is required'):
|
||||
params.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_train_size_too_high(valid_train_params_dict):
|
||||
"""Test validate_business_rules raises error when train_size > 100."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
valid_train_params_dict['train_size'] = 101
|
||||
params = TrainModelParams.from_dict(valid_train_params_dict)
|
||||
|
||||
with pytest.raises(ValueError, match='train_size must be between 10 and 100, got 101'):
|
||||
def test_validate_business_rules_train_size_out_of_range(valid_train_params_dict):
|
||||
d = copy.deepcopy(valid_train_params_dict)
|
||||
d['train_size'] = 5
|
||||
params = TrainModelParams.from_dict(d)
|
||||
with pytest.raises(ValueError, match='train_size must be between'):
|
||||
params.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_empty_variable_columns(valid_train_params_dict):
|
||||
"""Test validate_business_rules raises error when variable_columns is empty."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
valid_train_params_dict['variable_columns'] = []
|
||||
params = TrainModelParams.from_dict(valid_train_params_dict)
|
||||
|
||||
d = copy.deepcopy(valid_train_params_dict)
|
||||
d['variable_columns'] = []
|
||||
params = TrainModelParams.from_dict(d)
|
||||
with pytest.raises(ValueError, match='variable_columns cannot be empty'):
|
||||
params.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_negative_lag_train(valid_train_params_dict):
|
||||
"""Test validate_business_rules raises error when lag_train is negative."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
valid_train_params_dict['lag_train'] = {'var1': -1, 'var2': 5}
|
||||
params = TrainModelParams.from_dict(valid_train_params_dict)
|
||||
|
||||
with pytest.raises(ValueError, match='lag_train for var1 must be non-negative, got -1'):
|
||||
def test_validate_business_rules_empty_target(valid_train_params_dict):
|
||||
d = copy.deepcopy(valid_train_params_dict)
|
||||
d['target_variable'] = ' '
|
||||
params = TrainModelParams.from_dict(d)
|
||||
with pytest.raises(ValueError, match='target_variable cannot be empty'):
|
||||
params.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_negative_lag_val(valid_train_params_dict):
|
||||
"""Test validate_business_rules raises error when lag_val is negative."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
def test_from_dict_missing_required_key(valid_train_params_dict):
|
||||
d = copy.deepcopy(valid_train_params_dict)
|
||||
del d['bucket_name']
|
||||
with pytest.raises(ValueError, match='bucket_name is required'):
|
||||
TrainModelParams.from_dict(d)
|
||||
|
||||
valid_train_params_dict['lag_val'] = {'var1': 3, 'var2': -2}
|
||||
|
||||
def test_to_dict_roundtrip_keys(valid_train_params_dict):
|
||||
params = TrainModelParams.from_dict(valid_train_params_dict)
|
||||
|
||||
with pytest.raises(ValueError, match='lag_val for var2 must be non-negative, got -2'):
|
||||
params.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_negative_window(valid_train_params_dict):
|
||||
"""Test validate_business_rules raises error when window is negative."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
valid_train_params_dict['window'] = -5
|
||||
params = TrainModelParams.from_dict(valid_train_params_dict)
|
||||
|
||||
with pytest.raises(ValueError, match='window must be non-negative, got -5'):
|
||||
params.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_mismatched_limit_keys(valid_train_params_dict):
|
||||
"""Test validate_business_rules raises error when low_lim and upp_lim keys don't match."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
valid_train_params_dict['low_lim'] = {'var1': 0.0}
|
||||
valid_train_params_dict['upp_lim'] = {'var1': 10.0, 'var2': 20.0}
|
||||
params = TrainModelParams.from_dict(valid_train_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_train_params_dict):
|
||||
"""Test validate_business_rules raises error when low_lim >= upp_lim."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
valid_train_params_dict['low_lim'] = {'var1': 15.0, 'var2': 1.0}
|
||||
valid_train_params_dict['upp_lim'] = {'var1': 10.0, 'var2': 20.0}
|
||||
params = TrainModelParams.from_dict(valid_train_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_train_params_dict):
|
||||
"""Test validate_business_rules raises error when low_lim == upp_lim."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
valid_train_params_dict['low_lim'] = {'var1': 10.0, 'var2': 1.0}
|
||||
valid_train_params_dict['upp_lim'] = {'var1': 10.0, 'var2': 20.0}
|
||||
params = TrainModelParams.from_dict(valid_train_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_empty_bucket_name(valid_train_params_dict):
|
||||
"""Test validate_business_rules raises error when bucket_name is empty."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
valid_train_params_dict['bucket_name'] = ''
|
||||
params = TrainModelParams.from_dict(valid_train_params_dict)
|
||||
|
||||
with pytest.raises(ValueError, match='bucket_name cannot be empty or whitespace'):
|
||||
params.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_whitespace_bucket_name(valid_train_params_dict):
|
||||
"""Test validate_business_rules raises error when bucket_name is whitespace."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
valid_train_params_dict['bucket_name'] = ' '
|
||||
params = TrainModelParams.from_dict(valid_train_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_train_params_dict):
|
||||
"""Test validate_business_rules raises error when file_name is empty."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
valid_train_params_dict['file_name'] = ''
|
||||
params = TrainModelParams.from_dict(valid_train_params_dict)
|
||||
|
||||
with pytest.raises(ValueError, match='file_name cannot be empty or whitespace'):
|
||||
params.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_whitespace_file_name(valid_train_params_dict):
|
||||
"""Test validate_business_rules raises error when file_name is whitespace."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
valid_train_params_dict['file_name'] = ' \t '
|
||||
params = TrainModelParams.from_dict(valid_train_params_dict)
|
||||
|
||||
with pytest.raises(ValueError, match='file_name cannot be empty or whitespace'):
|
||||
params.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_model_name_empty(valid_train_params_dict):
|
||||
"""Test validate_business_rules raises error when model_name is empty."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
valid_train_params_dict['model_name'] = ''
|
||||
params = TrainModelParams.from_dict(valid_train_params_dict)
|
||||
|
||||
with pytest.raises(ValueError, match='model_name cannot be empty or whitespace'):
|
||||
params.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_train_size_boundary_10(valid_train_params_dict):
|
||||
"""Test validate_business_rules accepts train_size = 10 (lower boundary)."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
valid_train_params_dict['train_size'] = 10
|
||||
params = TrainModelParams.from_dict(valid_train_params_dict)
|
||||
|
||||
params.validate_business_rules() # Should not raise
|
||||
|
||||
|
||||
def test_validate_business_rules_train_size_boundary_100(valid_train_params_dict):
|
||||
"""Test validate_business_rules accepts train_size = 100 (upper boundary)."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
valid_train_params_dict['train_size'] = 100
|
||||
params = TrainModelParams.from_dict(valid_train_params_dict)
|
||||
|
||||
params.validate_business_rules() # Should not raise
|
||||
|
||||
|
||||
def test_validate_business_rules_zero_lag_train(valid_train_params_dict):
|
||||
"""Test validate_business_rules accepts lag_train = 0."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
valid_train_params_dict['lag_train'] = {'var1': 0, 'var2': 0}
|
||||
params = TrainModelParams.from_dict(valid_train_params_dict)
|
||||
|
||||
params.validate_business_rules() # Should not raise
|
||||
|
||||
|
||||
def test_validate_business_rules_zero_lag_val(valid_train_params_dict):
|
||||
"""Test validate_business_rules accepts lag_val = 0."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
valid_train_params_dict['lag_val'] = {'var1': 0, 'var2': 0}
|
||||
params = TrainModelParams.from_dict(valid_train_params_dict)
|
||||
|
||||
params.validate_business_rules() # Should not raise
|
||||
|
||||
|
||||
def test_validate_business_rules_zero_window(valid_train_params_dict):
|
||||
"""Test validate_business_rules accepts window = 0."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
valid_train_params_dict['window'] = 0
|
||||
params = TrainModelParams.from_dict(valid_train_params_dict)
|
||||
|
||||
params.validate_business_rules() # Should not raise
|
||||
|
||||
|
||||
def test_validate_business_rules_empty_limits(valid_train_params_dict):
|
||||
"""Test validate_business_rules accepts empty low_lim and upp_lim."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
valid_train_params_dict['low_lim'] = {}
|
||||
valid_train_params_dict['upp_lim'] = {}
|
||||
params = TrainModelParams.from_dict(valid_train_params_dict)
|
||||
|
||||
params.validate_business_rules() # Should not raise
|
||||
|
||||
|
||||
# ============================================================================
|
||||
# Additional tests for 100% coverage
|
||||
# ============================================================================
|
||||
|
||||
|
||||
def test_validate_business_rules_degree_less_than_1(valid_train_params_dict):
|
||||
"""Test validate_business_rules raises error when degree < 1."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
valid_train_params_dict['degree'] = 0
|
||||
params = TrainModelParams.from_dict(valid_train_params_dict)
|
||||
|
||||
with pytest.raises(ValueError, match='degree must be at least 1, got 0'):
|
||||
params.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_invalid_nan_treatment(valid_train_params_dict):
|
||||
"""Test validate_business_rules raises error for invalid nan_treatment."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
valid_train_params_dict['nan_treatment'] = 'invalid_treatment'
|
||||
params = TrainModelParams.from_dict(valid_train_params_dict)
|
||||
|
||||
with pytest.raises(ValueError, match='nan_treatment must be one of'):
|
||||
params.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_invalid_scaler_name(valid_train_params_dict):
|
||||
"""Test validate_business_rules raises error for invalid scaler_name."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
valid_train_params_dict['scaler_name'] = 'Invalid Scaler'
|
||||
params = TrainModelParams.from_dict(valid_train_params_dict)
|
||||
|
||||
with pytest.raises(ValueError, match='scaler_name must be one of'):
|
||||
params.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_invalid_model_name(valid_train_params_dict):
|
||||
"""Test validate_business_rules raises error for invalid model_name."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
valid_train_params_dict['model_name'] = 'Invalid Model'
|
||||
params = TrainModelParams.from_dict(valid_train_params_dict)
|
||||
|
||||
with pytest.raises(ValueError, match='model_name must be one of'):
|
||||
params.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_polynomial_regression_degree_less_than_2(valid_train_params_dict):
|
||||
"""Test validate_business_rules raises error for Polynomial Regression with degree < 2."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
valid_train_params_dict['model_name'] = 'Polynomial Regression'
|
||||
valid_train_params_dict['degree'] = 1
|
||||
params = TrainModelParams.from_dict(valid_train_params_dict)
|
||||
|
||||
with pytest.raises(ValueError, match='degree must be at least 2 for Polynomial Regression'):
|
||||
params.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_polynomial_regression_without_scaler(valid_train_params_dict):
|
||||
"""Test validate_business_rules raises error for Polynomial Regression without scaler."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
valid_train_params_dict['model_name'] = 'Polynomial Regression'
|
||||
valid_train_params_dict['degree'] = 2
|
||||
valid_train_params_dict['scaler_name'] = 'None'
|
||||
params = TrainModelParams.from_dict(valid_train_params_dict)
|
||||
|
||||
with pytest.raises(ValueError, match='scaler_name must be set'):
|
||||
params.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_linear_regression_degree_not_1(valid_train_params_dict):
|
||||
"""Test validate_business_rules raises error for Linear Regression with degree != 1."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
valid_train_params_dict['model_name'] = 'Linear Regression'
|
||||
valid_train_params_dict['degree'] = 2
|
||||
params = TrainModelParams.from_dict(valid_train_params_dict)
|
||||
|
||||
with pytest.raises(ValueError, match='degree must be 1 for Linear Regression, got 2'):
|
||||
params.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_removed_intervals_not_list(valid_train_params_dict):
|
||||
"""Test validate_business_rules raises error when removed_intervals item is not list."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
valid_train_params_dict['removed_intervals'] = ['not_a_list']
|
||||
params = TrainModelParams.from_dict(valid_train_params_dict)
|
||||
|
||||
with pytest.raises(ValueError, match='removed_intervals\\[0\\] must be a list or tuple'):
|
||||
params.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_removed_intervals_too_short(valid_train_params_dict):
|
||||
"""Test validate_business_rules raises error when removed_intervals item has < 2 elements."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
valid_train_params_dict['removed_intervals'] = [['only_one_element']]
|
||||
params = TrainModelParams.from_dict(valid_train_params_dict)
|
||||
|
||||
with pytest.raises(ValueError, match='removed_intervals\\[0\\] must have at least 2 elements'):
|
||||
params.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_start_date_wrong_type(valid_train_params_dict):
|
||||
"""Test validate_business_rules raises error when start_date is not a string."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
# Create params normally first, then modify start_date to bypass from_dict validation
|
||||
params = TrainModelParams.from_dict(valid_train_params_dict)
|
||||
params.start_date = 12345 # type: ignore
|
||||
|
||||
with pytest.raises(TypeError, match='start_date must be a string, got int'):
|
||||
params.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_end_date_wrong_type(valid_train_params_dict):
|
||||
"""Test validate_business_rules raises error when end_date is not a string."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
# Create params normally first, then modify end_date to bypass from_dict validation
|
||||
params = TrainModelParams.from_dict(valid_train_params_dict)
|
||||
params.end_date = 12345 # type: ignore
|
||||
|
||||
with pytest.raises(TypeError, match='end_date must be a string, got int'):
|
||||
params.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_empty_target_variable(valid_train_params_dict):
|
||||
"""Test validate_business_rules raises error when target_variable is empty."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
valid_train_params_dict['target_variable'] = ''
|
||||
params = TrainModelParams.from_dict(valid_train_params_dict)
|
||||
|
||||
with pytest.raises(ValueError, match='target_variable cannot be empty or whitespace'):
|
||||
params.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_whitespace_target_variable(valid_train_params_dict):
|
||||
"""Test validate_business_rules raises error when target_variable is whitespace."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
valid_train_params_dict['target_variable'] = ' '
|
||||
params = TrainModelParams.from_dict(valid_train_params_dict)
|
||||
|
||||
with pytest.raises(ValueError, match='target_variable cannot be empty or whitespace'):
|
||||
params.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_valid_removed_intervals(valid_train_params_dict):
|
||||
"""Test validate_business_rules accepts valid removed_intervals."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
valid_train_params_dict['removed_intervals'] = [['2023-01-01', '2023-01-02']]
|
||||
params = TrainModelParams.from_dict(valid_train_params_dict)
|
||||
|
||||
params.validate_business_rules() # Should not raise
|
||||
|
||||
|
||||
def test_validate_business_rules_valid_start_and_end_date(valid_train_params_dict):
|
||||
"""Test validate_business_rules accepts valid start_date and end_date."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
valid_train_params_dict['start_date'] = '2023-01-01'
|
||||
valid_train_params_dict['end_date'] = '2023-12-31'
|
||||
params = TrainModelParams.from_dict(valid_train_params_dict)
|
||||
|
||||
params.validate_business_rules() # Should not raise
|
||||
|
||||
|
||||
def test_validate_business_rules_invalid_date_format(valid_train_params_dict):
|
||||
"""Test validate_business_rules raises when date_format is not allowed."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
params = TrainModelParams.from_dict(valid_train_params_dict)
|
||||
params.date_format = 'yyyy-MM-dd'
|
||||
|
||||
d = params.to_dict()
|
||||
assert 'variable_columns' in d
|
||||
assert d['experiment_run_id'] == 1
|
||||
|
||||
|
||||
def test_coerce_experiment_run_id_float():
|
||||
assert TrainModelParams._coerce_experiment_run_id(2.0) == 2
|
||||
|
||||
|
||||
def test_coerce_experiment_run_id_none_raises():
|
||||
with pytest.raises(ValueError, match='experiment_run_id is required'):
|
||||
TrainModelParams._coerce_experiment_run_id(None)
|
||||
|
||||
|
||||
def test_coerce_experiment_run_id_invalid_type():
|
||||
with pytest.raises(TypeError, match='integer or numeric string'):
|
||||
TrainModelParams._coerce_experiment_run_id([1])
|
||||
|
||||
|
||||
def test_validate_model_param_schema_validation_error(valid_train_params_dict):
|
||||
d = copy.deepcopy(valid_train_params_dict)
|
||||
d['model_metadata'] = {
|
||||
'schemas': {
|
||||
'components': {
|
||||
'schemas': {
|
||||
'data_model': {'type': 'object', 'properties': {'x': {'type': 'integer'}}, 'required': ['x']},
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
p = TrainModelParams.from_dict(d)
|
||||
p.data_model_kwargs = {}
|
||||
with pytest.raises(ValueError, match='Model parameters validation failed'):
|
||||
p.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_model_param_unexpected_validator_error(valid_train_params_dict):
|
||||
d = copy.deepcopy(valid_train_params_dict)
|
||||
d['model_metadata'] = {
|
||||
'schemas': {
|
||||
'components': {
|
||||
'schemas': {
|
||||
'data_model': {'type': 'object'},
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
p = TrainModelParams.from_dict(d)
|
||||
with patch('model_manager.utils.models.train_model_params.Draft202012Validator') as m:
|
||||
m.return_value.validate.side_effect = RuntimeError('boom')
|
||||
with pytest.raises(ValueError, match='Unexpected error'):
|
||||
p.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_date_format_invalid(valid_train_params_dict):
|
||||
d = copy.deepcopy(valid_train_params_dict)
|
||||
d['date_format'] = 'not-an-allowed-format'
|
||||
p = TrainModelParams.from_dict(d)
|
||||
with pytest.raises(ValueError, match='Invalid date_format'):
|
||||
params.validate_business_rules()
|
||||
p.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_valid_date_format(valid_train_params_dict):
|
||||
"""Test validate_business_rules accepts allowed date_format."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
params = TrainModelParams.from_dict(valid_train_params_dict)
|
||||
params.date_format = 'yyyy-MM-dd HH:mm:ss'
|
||||
|
||||
params.validate_business_rules() # Should not raise
|
||||
def test_validate_required_strings_whitespace_bucket_file_model(valid_train_params_dict):
|
||||
for field, msg in [
|
||||
('bucket_name', 'bucket_name cannot be empty'),
|
||||
('file_name', 'file_name cannot be empty'),
|
||||
('model_name', 'model_name cannot be empty'),
|
||||
]:
|
||||
d = copy.deepcopy(valid_train_params_dict)
|
||||
d[field] = ' '
|
||||
p = TrainModelParams.from_dict(d)
|
||||
with pytest.raises(ValueError, match=msg):
|
||||
p.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_polynomial_regression_valid(valid_train_params_dict):
|
||||
"""Test validate_business_rules accepts valid Polynomial Regression config."""
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
|
||||
valid_train_params_dict['model_name'] = 'Polynomial Regression'
|
||||
valid_train_params_dict['degree'] = 2
|
||||
valid_train_params_dict['scaler_name'] = 'Standard Scaler'
|
||||
params = TrainModelParams.from_dict(valid_train_params_dict)
|
||||
|
||||
params.validate_business_rules() # Should not raise
|
||||
def test_validate_model_param_only_data_model_schema(valid_train_params_dict):
|
||||
d = copy.deepcopy(valid_train_params_dict)
|
||||
d['model_metadata'] = {
|
||||
'schemas': {'components': {'schemas': {'data_model': {'type': 'object'}}}}
|
||||
}
|
||||
p = TrainModelParams.from_dict(d)
|
||||
p.data_model_kwargs = {}
|
||||
p.validate_business_rules()
|
||||
|
||||
|
||||
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_model_param_only_model_schema(valid_train_params_dict):
|
||||
d = copy.deepcopy(valid_train_params_dict)
|
||||
d['model_metadata'] = {
|
||||
'schemas': {'components': {'schemas': {'model': {'type': 'object'}}}}
|
||||
}
|
||||
p = TrainModelParams.from_dict(d)
|
||||
p.model_kwargs = {}
|
||||
p.validate_business_rules()
|
||||
|
||||
|
||||
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_model_param_only_opt_params_schema(valid_train_params_dict):
|
||||
d = copy.deepcopy(valid_train_params_dict)
|
||||
d['model_metadata'] = {
|
||||
'schemas': {'components': {'schemas': {'opt_params': {'type': 'object'}}}}
|
||||
}
|
||||
p = TrainModelParams.from_dict(d)
|
||||
p.opt_params = {}
|
||||
p.validate_business_rules()
|
||||
|
||||
|
||||
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)
|
||||
def test_validate_model_param_all_schema_branches(valid_train_params_dict):
|
||||
d = copy.deepcopy(valid_train_params_dict)
|
||||
d['model_metadata'] = {
|
||||
'schemas': {
|
||||
'components': {
|
||||
'schemas': {
|
||||
'data_model': {'type': 'object'},
|
||||
'model': {'type': 'object'},
|
||||
'opt_params': {'type': 'object'},
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
p = TrainModelParams.from_dict(d)
|
||||
p.data_model_kwargs = {}
|
||||
p.model_kwargs = {}
|
||||
p.opt_params = {}
|
||||
p.validate_business_rules()
|
||||
|
||||
@@ -1,7 +1,5 @@
|
||||
"""Unit tests for TrainModelResult dataclass."""
|
||||
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pandas as pd
|
||||
import pytest
|
||||
|
||||
@@ -10,242 +8,64 @@ from model_manager.utils.models.train_model_result import TrainModelResult
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sample_params():
|
||||
"""Create sample TrainModelParams for testing."""
|
||||
return TrainModelParams(
|
||||
variable_columns=['var1', 'var2'],
|
||||
lag_train={'var1': 5, 'var2': 5},
|
||||
lag_val={'var1': 3, 'var2': 3},
|
||||
target_variable='target',
|
||||
rem_static_win=True,
|
||||
low_lim={'var1': 0.0, 'var2': 0.0},
|
||||
upp_lim={'var1': 100.0, 'var2': 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,
|
||||
removed_intervals=[],
|
||||
model_name='Linear Regression',
|
||||
degree=1,
|
||||
interaction_only=False,
|
||||
nan_treatment='drop',
|
||||
start_date=None,
|
||||
end_date=None,
|
||||
scaler_name='Standard Scaler',
|
||||
support_filters={},
|
||||
static_threshold=None,
|
||||
date_column=None,
|
||||
date_format=None,
|
||||
def sample_params() -> TrainModelParams:
|
||||
"""Minimal TrainModelParams for TrainModelResult tests."""
|
||||
return TrainModelParams.from_dict(
|
||||
{
|
||||
'variable_columns': ['a'],
|
||||
'target_variable': 't',
|
||||
'bucket_name': 'b',
|
||||
'file_name': 'f.csv',
|
||||
'line_separator': '\n',
|
||||
'decimal_separator': '.',
|
||||
'date_column': None,
|
||||
'date_format': None,
|
||||
'train_size': 80,
|
||||
'shuffle': True,
|
||||
'random_state': 42,
|
||||
'experiment_run_id': 1,
|
||||
'model_name': 'Linear Regression',
|
||||
'val_file_name': None,
|
||||
'data_model_kwargs': {},
|
||||
'model_kwargs': {},
|
||||
'opt_params': {},
|
||||
'model_type': 'linear_regression',
|
||||
'model_id': None,
|
||||
'model_metadata': {'schemas': {'components': {'schemas': {}}}},
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sample_dataframes():
|
||||
"""Create sample DataFrames for testing."""
|
||||
X_train = pd.DataFrame({'var1': [1, 2, 3], 'var2': [4, 5, 6]})
|
||||
X_test = pd.DataFrame({'var1': [7, 8], 'var2': [9, 10]})
|
||||
y_train = pd.DataFrame({'target': [10, 20, 30]})
|
||||
y_test = pd.DataFrame({'target': [40, 50]})
|
||||
return X_train, X_test, y_train, y_test
|
||||
def sample_frames():
|
||||
train = pd.DataFrame({'a': [1, 2], 't': [1.0, 2.0]})
|
||||
val = pd.DataFrame({'a': [3], 't': [3.0]})
|
||||
return train, val
|
||||
|
||||
|
||||
def test_train_model_result_creation(sample_params, sample_dataframes):
|
||||
"""Test creating TrainModelResult with required fields."""
|
||||
x_train, x_test, y_train, y_test = sample_dataframes
|
||||
process_data = MagicMock()
|
||||
regr = MagicMock()
|
||||
|
||||
result = TrainModelResult(
|
||||
params=sample_params,
|
||||
process_data=process_data,
|
||||
x_train=x_train,
|
||||
x_test=x_test,
|
||||
y_train=y_train,
|
||||
y_test=y_test,
|
||||
regr=regr,
|
||||
)
|
||||
|
||||
assert result.params == sample_params
|
||||
assert result.process_data == process_data
|
||||
assert result.x_train.equals(x_train)
|
||||
assert result.x_test.equals(x_test)
|
||||
assert result.y_train.equals(y_train)
|
||||
assert result.y_test.equals(y_test)
|
||||
assert result.regr == regr
|
||||
assert result.scaler_dict == scaler_dict
|
||||
|
||||
|
||||
def test_train_model_result_optional_fields_default_none(sample_params, sample_dataframes):
|
||||
"""Test that optional fields default to None."""
|
||||
x_train, x_test, y_train, y_test = sample_dataframes
|
||||
|
||||
result = TrainModelResult(
|
||||
params=sample_params,
|
||||
process_data=MagicMock(),
|
||||
x_train=x_train,
|
||||
x_test=x_test,
|
||||
y_train=y_train,
|
||||
y_test=y_test,
|
||||
regr=MagicMock(),
|
||||
)
|
||||
|
||||
assert result.y_pred is None
|
||||
assert result.mse_val is None
|
||||
assert result.mae_val is None
|
||||
assert result.r2_val is None
|
||||
def test_train_model_result_creation(sample_params, sample_frames):
|
||||
train, val = sample_frames
|
||||
result = TrainModelResult(params=sample_params, train_data=train, val_data=val)
|
||||
assert result.params is sample_params
|
||||
assert result.train_data.equals(train)
|
||||
assert result.val_data.equals(val)
|
||||
assert result.run_name is None
|
||||
assert result.report_path is None
|
||||
assert result.train_data_path is None
|
||||
assert result.test_data_path is None
|
||||
assert result.run_dir is None
|
||||
|
||||
|
||||
def test_train_model_result_with_metrics(sample_params, sample_dataframes):
|
||||
"""Test TrainModelResult with metrics populated."""
|
||||
x_train, x_test, y_train, y_test = sample_dataframes
|
||||
y_pred = pd.Series([41, 49])
|
||||
|
||||
def test_train_model_result_optional_paths(sample_params, sample_frames):
|
||||
train, val = sample_frames
|
||||
result = TrainModelResult(
|
||||
params=sample_params,
|
||||
process_data=MagicMock(),
|
||||
x_train=x_train,
|
||||
x_test=x_test,
|
||||
y_train=y_train,
|
||||
y_test=y_test,
|
||||
regr=MagicMock(),
|
||||
y_pred=y_pred,
|
||||
mse_val=1.5,
|
||||
mae_val=1.2,
|
||||
r2_val=0.95,
|
||||
train_data=train,
|
||||
val_data=val,
|
||||
run_name='run-1',
|
||||
run_id='rid',
|
||||
run_dir='/tmp/x',
|
||||
mse_val=0.1,
|
||||
mae_val=0.2,
|
||||
r2_val=0.99,
|
||||
)
|
||||
|
||||
assert result.y_pred.equals(y_pred)
|
||||
assert result.mse_val == 1.5
|
||||
assert result.mae_val == 1.2
|
||||
assert result.r2_val == 0.95
|
||||
|
||||
|
||||
def test_train_model_result_with_artifact_paths(sample_params, sample_dataframes):
|
||||
"""Test TrainModelResult with artifact paths populated."""
|
||||
x_train, x_test, y_train, y_test = sample_dataframes
|
||||
|
||||
result = TrainModelResult(
|
||||
params=sample_params,
|
||||
process_data=MagicMock(),
|
||||
x_train=x_train,
|
||||
x_test=x_test,
|
||||
y_train=y_train,
|
||||
y_test=y_test,
|
||||
regr=MagicMock(),
|
||||
run_name='test-experiment-1',
|
||||
report_path='/path/to/report.html',
|
||||
train_data_path='/path/to/train_data.csv',
|
||||
test_data_path='/path/to/test_data.csv',
|
||||
run_dir='/path/to/run_dir',
|
||||
)
|
||||
|
||||
assert result.run_name == 'test-experiment-1'
|
||||
assert result.report_path == '/path/to/report.html'
|
||||
assert result.train_data_path == '/path/to/train_data.csv'
|
||||
assert result.test_data_path == '/path/to/test_data.csv'
|
||||
assert result.run_dir == '/path/to/run_dir'
|
||||
|
||||
|
||||
def test_train_model_result_is_dataclass(sample_params, sample_dataframes):
|
||||
"""Test that TrainModelResult is a dataclass."""
|
||||
x_train, x_test, y_train, y_test = sample_dataframes
|
||||
|
||||
result = TrainModelResult(
|
||||
params=sample_params,
|
||||
process_data=MagicMock(),
|
||||
x_train=x_train,
|
||||
x_test=x_test,
|
||||
y_train=y_train,
|
||||
y_test=y_test,
|
||||
regr=MagicMock(),
|
||||
scaler_dict={},
|
||||
)
|
||||
|
||||
# Dataclasses have __dataclass_fields__ attribute
|
||||
assert hasattr(result, '__dataclass_fields__')
|
||||
assert 'params' in result.__dataclass_fields__
|
||||
assert 'process_data' in result.__dataclass_fields__
|
||||
assert 'x_train' in result.__dataclass_fields__
|
||||
|
||||
|
||||
def test_train_model_result_field_count():
|
||||
"""Test that TrainModelResult has exactly 19 fields."""
|
||||
from dataclasses import fields
|
||||
|
||||
result_fields = fields(TrainModelResult)
|
||||
assert len(result_fields) == 19
|
||||
|
||||
field_names = {f.name for f in result_fields}
|
||||
expected_fields = {
|
||||
'params',
|
||||
'process_data',
|
||||
'x_train',
|
||||
'x_test',
|
||||
'y_train',
|
||||
'y_test',
|
||||
'regr',
|
||||
'y_pred',
|
||||
'y_train_pred',
|
||||
'mse_val',
|
||||
'mae_val',
|
||||
'r2_val',
|
||||
'equation',
|
||||
'equation_path',
|
||||
'run_name',
|
||||
'report_path',
|
||||
'train_data_path',
|
||||
'test_data_path',
|
||||
'run_dir',
|
||||
}
|
||||
assert field_names == expected_fields
|
||||
|
||||
|
||||
def test_train_model_result_complete_workflow(sample_params, sample_dataframes):
|
||||
"""Test TrainModelResult through a complete workflow simulation."""
|
||||
x_train, x_test, y_train, y_test = sample_dataframes
|
||||
|
||||
# Step 1: Create result after training
|
||||
result = TrainModelResult(
|
||||
params=sample_params,
|
||||
process_data=MagicMock(),
|
||||
x_train=x_train,
|
||||
x_test=x_test,
|
||||
y_train=y_train,
|
||||
y_test=y_test,
|
||||
regr=MagicMock(),
|
||||
)
|
||||
|
||||
# Step 2: Add predictions and metrics
|
||||
result.y_pred = pd.Series([41, 49])
|
||||
result.mse_val = 1.5
|
||||
result.mae_val = 1.2
|
||||
result.r2_val = 0.95
|
||||
|
||||
# Step 3: Add artifact paths
|
||||
result.run_name = 'test-experiment-1'
|
||||
result.report_path = '/path/to/report.html'
|
||||
result.train_data_path = '/path/to/train_data.csv'
|
||||
result.test_data_path = '/path/to/test_data.csv'
|
||||
result.run_dir = '/path/to/run_dir'
|
||||
|
||||
# Verify all fields are populated
|
||||
assert result.y_pred is not None
|
||||
assert result.mse_val == 1.5
|
||||
assert result.mae_val == 1.2
|
||||
assert result.r2_val == 0.95
|
||||
assert result.run_name == 'test-experiment-1'
|
||||
assert result.report_path == '/path/to/report.html'
|
||||
assert result.train_data_path == '/path/to/train_data.csv'
|
||||
assert result.test_data_path == '/path/to/test_data.csv'
|
||||
assert result.run_dir == '/path/to/run_dir'
|
||||
assert result.run_name == 'run-1'
|
||||
assert result.run_id == 'rid'
|
||||
assert result.run_dir == '/tmp/x'
|
||||
assert result.mse_val == 0.1
|
||||
|
||||
Reference in New Issue
Block a user