feat: require date_column in training parameters and update documentation
- Made `date_column` a required field in `TrainModelParams`, ensuring it must be present in the input data. - Updated related documentation in `input-sample.md`, `README.md`, and various test scenarios to reflect the change in requirement. - Adjusted the handling of `date_format` to default to `yyyy-MM-dd HH:mm:ss` if omitted, enhancing usability. - Refined test scenarios to include new examples and ensure compliance with the updated parameter structure. These changes improve the robustness of the model training workflow and clarify the expectations for input data.
This commit is contained in:
@@ -6,6 +6,7 @@ from unittest.mock import patch
|
||||
import pytest
|
||||
|
||||
from model_manager.utils.models.train_model_params import (
|
||||
DEFAULT_TRAIN_DATE_FORMAT,
|
||||
TrainModelParams,
|
||||
validate_frontend_date_format,
|
||||
)
|
||||
@@ -27,8 +28,8 @@ def valid_train_params_dict(minimal_model_metadata) -> dict:
|
||||
'file_name': 'test-file.csv',
|
||||
'line_separator': ',',
|
||||
'decimal_separator': '.',
|
||||
'date_column': None,
|
||||
'date_format': None,
|
||||
'date_column': 'timestamp',
|
||||
'date_format': 'yyyy-MM-dd HH:mm:ss',
|
||||
'train_size': 80,
|
||||
'shuffle': True,
|
||||
'random_state': 42,
|
||||
@@ -52,10 +53,47 @@ def test_from_dict_success(valid_train_params_dict):
|
||||
assert params.target_variable == 'target'
|
||||
assert params.bucket_name == 'test-bucket'
|
||||
assert params.experiment_run_id == 1
|
||||
assert params.experiment_name == 'Linear Regression_experiment'
|
||||
assert params.experiment_name == 'Linear Regression'
|
||||
assert params.model_metadata is valid_train_params_dict['model_metadata']
|
||||
|
||||
|
||||
def test_from_dict_date_format_omitted_uses_default(valid_train_params_dict):
|
||||
"""Missing date_format defaults to DEFAULT_TRAIN_DATE_FORMAT."""
|
||||
d = copy.deepcopy(valid_train_params_dict)
|
||||
del d['date_format']
|
||||
params = TrainModelParams.from_dict(d)
|
||||
assert params.date_format == DEFAULT_TRAIN_DATE_FORMAT
|
||||
|
||||
|
||||
def test_from_dict_date_format_blank_uses_default(valid_train_params_dict):
|
||||
d = copy.deepcopy(valid_train_params_dict)
|
||||
d['date_format'] = ' '
|
||||
params = TrainModelParams.from_dict(d)
|
||||
assert params.date_format == DEFAULT_TRAIN_DATE_FORMAT
|
||||
|
||||
|
||||
def test_from_dict_superfluous_date_column_camel_key_is_ignored(valid_train_params_dict):
|
||||
"""Only snake_case keys are read; dateColumn does not populate date_column."""
|
||||
d = copy.deepcopy(valid_train_params_dict)
|
||||
d['dateColumn'] = 'wrong_name'
|
||||
params = TrainModelParams.from_dict(d)
|
||||
assert params.date_column == 'timestamp'
|
||||
|
||||
|
||||
def test_from_dict_missing_date_column_raises(valid_train_params_dict):
|
||||
d = copy.deepcopy(valid_train_params_dict)
|
||||
del d['date_column']
|
||||
with pytest.raises(ValueError, match='date_column is required'):
|
||||
TrainModelParams.from_dict(d)
|
||||
|
||||
|
||||
def test_from_dict_date_format_non_string_raises(valid_train_params_dict):
|
||||
d = copy.deepcopy(valid_train_params_dict)
|
||||
d['date_format'] = 12345
|
||||
with pytest.raises(TypeError, match='date_format must be a string'):
|
||||
TrainModelParams.from_dict(d)
|
||||
|
||||
|
||||
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)
|
||||
@@ -131,6 +169,14 @@ def test_validate_business_rules_empty_target(valid_train_params_dict):
|
||||
params.validate_business_rules()
|
||||
|
||||
|
||||
def test_validate_business_rules_whitespace_date_column(valid_train_params_dict):
|
||||
d = copy.deepcopy(valid_train_params_dict)
|
||||
d['date_column'] = ' '
|
||||
params = TrainModelParams.from_dict(d)
|
||||
with pytest.raises(ValueError, match='date_column cannot be empty'):
|
||||
params.validate_business_rules()
|
||||
|
||||
|
||||
def test_from_dict_missing_required_key(valid_train_params_dict):
|
||||
d = copy.deepcopy(valid_train_params_dict)
|
||||
del d['bucket_name']
|
||||
|
||||
Reference in New Issue
Block a user