SIENTIAPDE-1249: Refactor TrainModelParams to use dataclass and add from_dict method for validation, remove no-cache-dir from pip install in quality gate workflow, and rename X_train/X_test to x_train/x_test in TrainModelResult.
This commit is contained in:
@@ -58,12 +58,21 @@ def test_train_model_params_creation_with_valid_params(valid_params_dict):
|
||||
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(**valid_params_dict)
|
||||
TrainModelParams.from_dict(valid_params_dict)
|
||||
|
||||
|
||||
def test_train_model_params_variable_columns_wrong_type_raises_error(valid_params_dict):
|
||||
@@ -71,7 +80,7 @@ def test_train_model_params_variable_columns_wrong_type_raises_error(valid_param
|
||||
valid_params_dict['variable_columns'] = 'not a list'
|
||||
|
||||
with pytest.raises(TypeError, match='variable_columns must be of type list'):
|
||||
TrainModelParams(**valid_params_dict)
|
||||
TrainModelParams.from_dict(valid_params_dict)
|
||||
|
||||
|
||||
def test_train_model_params_lag_train_none_raises_error(valid_params_dict):
|
||||
@@ -79,7 +88,7 @@ def test_train_model_params_lag_train_none_raises_error(valid_params_dict):
|
||||
valid_params_dict['lag_train'] = None
|
||||
|
||||
with pytest.raises(ValueError, match='lag_train is required and cannot be None'):
|
||||
TrainModelParams(**valid_params_dict)
|
||||
TrainModelParams.from_dict(valid_params_dict)
|
||||
|
||||
|
||||
def test_train_model_params_lag_train_wrong_type_raises_error(valid_params_dict):
|
||||
@@ -87,7 +96,7 @@ def test_train_model_params_lag_train_wrong_type_raises_error(valid_params_dict)
|
||||
valid_params_dict['lag_train'] = '5'
|
||||
|
||||
with pytest.raises(TypeError, match='lag_train must be of type int'):
|
||||
TrainModelParams(**valid_params_dict)
|
||||
TrainModelParams.from_dict(valid_params_dict)
|
||||
|
||||
|
||||
def test_train_model_params_target_variable_none_raises_error(valid_params_dict):
|
||||
@@ -95,7 +104,7 @@ def test_train_model_params_target_variable_none_raises_error(valid_params_dict)
|
||||
valid_params_dict['target_variable'] = None
|
||||
|
||||
with pytest.raises(ValueError, match='target_variable is required and cannot be None'):
|
||||
TrainModelParams(**valid_params_dict)
|
||||
TrainModelParams.from_dict(valid_params_dict)
|
||||
|
||||
|
||||
def test_train_model_params_target_variable_wrong_type_raises_error(valid_params_dict):
|
||||
@@ -103,7 +112,7 @@ def test_train_model_params_target_variable_wrong_type_raises_error(valid_params
|
||||
valid_params_dict['target_variable'] = 123
|
||||
|
||||
with pytest.raises(TypeError, match='target_variable must be of type str'):
|
||||
TrainModelParams(**valid_params_dict)
|
||||
TrainModelParams.from_dict(valid_params_dict)
|
||||
|
||||
|
||||
def test_train_model_params_boolean_fields(valid_params_dict):
|
||||
@@ -111,11 +120,11 @@ def test_train_model_params_boolean_fields(valid_params_dict):
|
||||
# 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(**valid_params_dict)
|
||||
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(**valid_params_dict)
|
||||
TrainModelParams.from_dict(valid_params_dict)
|
||||
|
||||
|
||||
def test_train_model_params_dict_fields(valid_params_dict):
|
||||
@@ -123,12 +132,12 @@ def test_train_model_params_dict_fields(valid_params_dict):
|
||||
# Test low_lim
|
||||
valid_params_dict['low_lim'] = None
|
||||
with pytest.raises(ValueError, match='low_lim is required and cannot be None'):
|
||||
TrainModelParams(**valid_params_dict)
|
||||
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(**valid_params_dict)
|
||||
TrainModelParams.from_dict(valid_params_dict)
|
||||
|
||||
|
||||
def test_train_model_params_bucket_name_none_raises_error(valid_params_dict):
|
||||
@@ -136,7 +145,7 @@ def test_train_model_params_bucket_name_none_raises_error(valid_params_dict):
|
||||
valid_params_dict['bucket_name'] = None
|
||||
|
||||
with pytest.raises(ValueError, match='bucket_name is required and cannot be None'):
|
||||
TrainModelParams(**valid_params_dict)
|
||||
TrainModelParams.from_dict(valid_params_dict)
|
||||
|
||||
|
||||
def test_train_model_params_file_name_none_raises_error(valid_params_dict):
|
||||
@@ -144,7 +153,7 @@ def test_train_model_params_file_name_none_raises_error(valid_params_dict):
|
||||
valid_params_dict['file_name'] = None
|
||||
|
||||
with pytest.raises(ValueError, match='file_name is required and cannot be None'):
|
||||
TrainModelParams(**valid_params_dict)
|
||||
TrainModelParams.from_dict(valid_params_dict)
|
||||
|
||||
|
||||
def test_train_model_params_experiment_run_id_none_raises_error(valid_params_dict):
|
||||
@@ -152,7 +161,7 @@ def test_train_model_params_experiment_run_id_none_raises_error(valid_params_dic
|
||||
valid_params_dict['experiment_run_id'] = None
|
||||
|
||||
with pytest.raises(ValueError, match='experiment_run_id is required and cannot be None'):
|
||||
TrainModelParams(**valid_params_dict)
|
||||
TrainModelParams.from_dict(valid_params_dict)
|
||||
|
||||
|
||||
def test_train_model_params_experiment_name_none_raises_error(valid_params_dict):
|
||||
@@ -160,14 +169,14 @@ def test_train_model_params_experiment_name_none_raises_error(valid_params_dict)
|
||||
valid_params_dict['experiment_name'] = None
|
||||
|
||||
with pytest.raises(ValueError, match='experiment_name is required and cannot be None'):
|
||||
TrainModelParams(**valid_params_dict)
|
||||
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(**valid_params_dict)
|
||||
params = TrainModelParams.from_dict(valid_params_dict)
|
||||
assert params.removed_intervals is None
|
||||
|
||||
|
||||
@@ -176,7 +185,7 @@ def test_train_model_params_removed_intervals_wrong_type_raises_error(valid_para
|
||||
valid_params_dict['removed_intervals'] = 'not a list'
|
||||
|
||||
with pytest.raises(TypeError, match='removed_intervals must be of type list'):
|
||||
TrainModelParams(**valid_params_dict)
|
||||
TrainModelParams.from_dict(valid_params_dict)
|
||||
|
||||
|
||||
def test_train_model_params_removed_intervals_with_values(valid_params_dict):
|
||||
@@ -186,7 +195,7 @@ def test_train_model_params_removed_intervals_with_values(valid_params_dict):
|
||||
('2023-02-01', '2023-02-05'),
|
||||
]
|
||||
|
||||
params = TrainModelParams(**valid_params_dict)
|
||||
params = TrainModelParams.from_dict(valid_params_dict)
|
||||
assert len(params.removed_intervals) == 2
|
||||
assert params.removed_intervals[0] == ('2023-01-01', '2023-01-10')
|
||||
|
||||
|
||||
@@ -48,7 +48,7 @@ def sample_dataframes():
|
||||
|
||||
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
|
||||
x_train, x_test, y_train, y_test = sample_dataframes
|
||||
process_data = MagicMock()
|
||||
regr = MagicMock()
|
||||
scaler_dict = {'var1': {'min': 0, 'max': 100}}
|
||||
@@ -56,8 +56,8 @@ def test_train_model_result_creation(sample_params, sample_dataframes):
|
||||
result = TrainModelResult(
|
||||
params=sample_params,
|
||||
process_data=process_data,
|
||||
X_train=X_train,
|
||||
X_test=X_test,
|
||||
x_train=x_train,
|
||||
x_test=x_test,
|
||||
y_train=y_train,
|
||||
y_test=y_test,
|
||||
regr=regr,
|
||||
@@ -66,8 +66,8 @@ def test_train_model_result_creation(sample_params, sample_dataframes):
|
||||
|
||||
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.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
|
||||
@@ -76,13 +76,13 @@ def test_train_model_result_creation(sample_params, sample_dataframes):
|
||||
|
||||
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
|
||||
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,
|
||||
x_train=x_train,
|
||||
x_test=x_test,
|
||||
y_train=y_train,
|
||||
y_test=y_test,
|
||||
regr=MagicMock(),
|
||||
@@ -102,14 +102,14 @@ def test_train_model_result_optional_fields_default_none(sample_params, sample_d
|
||||
|
||||
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
|
||||
x_train, x_test, y_train, y_test = sample_dataframes
|
||||
y_pred = pd.Series([41, 49])
|
||||
|
||||
result = TrainModelResult(
|
||||
params=sample_params,
|
||||
process_data=MagicMock(),
|
||||
X_train=X_train,
|
||||
X_test=X_test,
|
||||
x_train=x_train,
|
||||
x_test=x_test,
|
||||
y_train=y_train,
|
||||
y_test=y_test,
|
||||
regr=MagicMock(),
|
||||
@@ -128,13 +128,13 @@ def test_train_model_result_with_metrics(sample_params, sample_dataframes):
|
||||
|
||||
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
|
||||
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,
|
||||
x_train=x_train,
|
||||
x_test=x_test,
|
||||
y_train=y_train,
|
||||
y_test=y_test,
|
||||
regr=MagicMock(),
|
||||
@@ -155,13 +155,13 @@ def test_train_model_result_with_artifact_paths(sample_params, sample_dataframes
|
||||
|
||||
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
|
||||
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,
|
||||
x_train=x_train,
|
||||
x_test=x_test,
|
||||
y_train=y_train,
|
||||
y_test=y_test,
|
||||
regr=MagicMock(),
|
||||
@@ -172,7 +172,7 @@ def test_train_model_result_is_dataclass(sample_params, sample_dataframes):
|
||||
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__
|
||||
assert 'x_train' in result.__dataclass_fields__
|
||||
|
||||
|
||||
def test_train_model_result_field_count():
|
||||
@@ -186,8 +186,8 @@ def test_train_model_result_field_count():
|
||||
expected_fields = {
|
||||
'params',
|
||||
'process_data',
|
||||
'X_train',
|
||||
'X_test',
|
||||
'x_train',
|
||||
'x_test',
|
||||
'y_train',
|
||||
'y_test',
|
||||
'regr',
|
||||
@@ -207,14 +207,14 @@ def test_train_model_result_field_count():
|
||||
|
||||
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
|
||||
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,
|
||||
x_train=x_train,
|
||||
x_test=x_test,
|
||||
y_train=y_train,
|
||||
y_test=y_test,
|
||||
regr=MagicMock(),
|
||||
|
||||
Reference in New Issue
Block a user