- 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.
368 lines
14 KiB
Python
368 lines
14 KiB
Python
"""Unit tests for Training activities."""
|
|
|
|
from contextlib import contextmanager
|
|
from unittest.mock import MagicMock, patch
|
|
|
|
import pandas as pd
|
|
import pytest
|
|
|
|
from model_manager.utils.models.train_model_params import TrainModelParams
|
|
from model_manager.utils.models.train_model_result import TrainModelResult
|
|
|
|
|
|
def _minimal_params_dict():
|
|
return {
|
|
'variable_columns': ['a'],
|
|
'target_variable': 't',
|
|
'bucket_name': 'b',
|
|
'file_name': 'f.csv',
|
|
'line_separator': '\n',
|
|
'decimal_separator': '.',
|
|
'date_column': 'timestamp',
|
|
'date_format': 'yyyy-MM-dd HH:mm:ss',
|
|
'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': None,
|
|
}
|
|
|
|
|
|
@pytest.fixture
|
|
def training():
|
|
from model_manager.activities.training import Training
|
|
|
|
return Training(
|
|
mlflow_repository=MagicMock(),
|
|
plugin_store=MagicMock(),
|
|
minio_repository=MagicMock(),
|
|
logger=MagicMock(),
|
|
notification_handler=MagicMock(),
|
|
metrics_controller=MagicMock(),
|
|
)
|
|
|
|
|
|
def test_load_model_metadata_success(training):
|
|
training.plugin_store.get_model_index = MagicMock(
|
|
return_value={'schemas': {'components': {'schemas': {}}}}
|
|
)
|
|
inp = {**_minimal_params_dict(), 'metadata': {'w': '1'}}
|
|
out = training.load_model_metadata(inp)
|
|
assert 'model_metadata' in out
|
|
assert out['model_metadata']['schemas']
|
|
|
|
|
|
def test_load_model_metadata_notifies_on_error(training):
|
|
training.plugin_store.get_model_index = MagicMock(side_effect=RuntimeError('idx'))
|
|
training.send_notification = MagicMock()
|
|
inp = {**_minimal_params_dict(), 'metadata': {}}
|
|
with pytest.raises(RuntimeError, match='idx'):
|
|
training.load_model_metadata(inp)
|
|
training.send_notification.assert_called_once()
|
|
|
|
|
|
def test_validate_train_params_success(training):
|
|
pdict = _minimal_params_dict()
|
|
pdict['model_metadata'] = {'schemas': {'components': {'schemas': {}}}}
|
|
inp = {**pdict, 'metadata': {}}
|
|
out = training.validate_train_params(inp)
|
|
assert isinstance(out, dict)
|
|
assert out['target_variable'] == 't'
|
|
|
|
|
|
def test_validate_train_params_notifies(training):
|
|
training.send_notification = MagicMock()
|
|
inp = {'metadata': {}, 'experiment_run_id': 1}
|
|
with pytest.raises((KeyError, ValueError, TypeError)):
|
|
training.validate_train_params(inp)
|
|
training.send_notification.assert_called_once()
|
|
|
|
|
|
def test_train_model_download_fails_notifies(training):
|
|
"""train_model notifies and re-raises when MinIO download fails."""
|
|
tp = TrainModelParams.from_dict(
|
|
{
|
|
**_minimal_params_dict(),
|
|
'model_metadata': {'schemas': {'components': {'schemas': {}}}},
|
|
}
|
|
)
|
|
training.minio_repository.download_file = MagicMock(side_effect=OSError('minio'))
|
|
training.send_notification = MagicMock()
|
|
with pytest.raises(OSError, match='minio'):
|
|
training.train_model({'metadata': {'pod': 'x'}, 'train_params': tp.to_dict()})
|
|
training.send_notification.assert_called_once()
|
|
|
|
|
|
def test_cleanup_resources(training):
|
|
training.data_manager_repository.cleanup_run_directory = MagicMock()
|
|
training.cleanup_resources({'metadata': {}, 'run_dir': '/tmp/x'})
|
|
training.data_manager_repository.cleanup_run_directory.assert_called_once_with('/tmp/x', {})
|
|
|
|
|
|
def test_cleanup_resources_notifies_on_error(training):
|
|
training.data_manager_repository.cleanup_run_directory = MagicMock(
|
|
side_effect=RuntimeError('rm')
|
|
)
|
|
training.send_notification = MagicMock()
|
|
with pytest.raises(RuntimeError, match='rm'):
|
|
training.cleanup_resources({'metadata': {'pod': 'p'}, 'run_dir': '/tmp/x'})
|
|
training.send_notification.assert_called_once()
|
|
|
|
|
|
@patch('model_manager.activities.training.mlflow')
|
|
def test_train_model_success_serializes_result(mock_mlflow, training):
|
|
"""Exercise train_model happy path with mocks (MinIO, plugin wrapper, MLflow)."""
|
|
tp = TrainModelParams.from_dict(
|
|
{
|
|
**_minimal_params_dict(),
|
|
'model_metadata': {'schemas': {'components': {'schemas': {}}}},
|
|
}
|
|
)
|
|
train_df = pd.DataFrame({'a': [1.0, 2.0], 't': [1.0, 2.0]})
|
|
val_df = pd.DataFrame({'a': [1.0], 't': [1.0]})
|
|
tmr = TrainModelResult(params=tp, train_data=train_df, val_data=val_df)
|
|
|
|
training.minio_repository.download_file = MagicMock(return_value=b'csv')
|
|
training.data_manager_repository.prepare_training_data = MagicMock(return_value=tmr)
|
|
training.data_manager_repository.compute_regression_metrics = MagicMock(
|
|
side_effect=lambda x, _w, **_kw: (
|
|
setattr(x, 'mse_val', 0.1),
|
|
setattr(x, 'mae_val', 0.2),
|
|
setattr(x, 'r2_val', 0.9),
|
|
x,
|
|
)[-1]
|
|
)
|
|
|
|
def _fill_report(x, **_kw):
|
|
x.report_path = '/tmp/report.html'
|
|
x.train_data_path = '/tmp/train.csv'
|
|
x.test_data_path = '/tmp/test.csv'
|
|
x.run_dir = '/tmp/run'
|
|
return x
|
|
|
|
training.data_manager_repository.generate_report = MagicMock(side_effect=_fill_report)
|
|
|
|
wrapper = MagicMock()
|
|
wrapper.transform = MagicMock(side_effect=[(train_df, None), (val_df, None)])
|
|
pred_train = pd.DataFrame({'p': [1.0, 2.0]})
|
|
pred_val = pd.DataFrame({'p': [1.0]})
|
|
wrapper.predict = MagicMock(side_effect=[(pred_train, None), (pred_val, None)])
|
|
wrapper.store_model = MagicMock()
|
|
training.plugin_store.get_model = MagicMock(return_value=wrapper)
|
|
|
|
@contextmanager
|
|
def _run_ctx(*_a, **_k):
|
|
info = MagicMock()
|
|
info.run_name = 'run-n'
|
|
info.run_id = 'run-i'
|
|
yield info
|
|
|
|
training.mlflow_repository.start_run = _run_ctx
|
|
|
|
out = training.train_model({'metadata': {'pod': 'p'}, 'train_params': tp.to_dict()})
|
|
assert out['run_name'] is None
|
|
assert out['run_id'] == 'run-i'
|
|
assert out['run_dir'] == '/tmp/run'
|
|
mock_mlflow.log_param.assert_any_call('mse_val', 0.1)
|
|
mock_mlflow.log_param.assert_any_call('mae_val', 0.2)
|
|
mock_mlflow.log_param.assert_any_call('r2_val', 0.9)
|
|
mock_mlflow.log_artifact.assert_called()
|
|
|
|
|
|
@patch('model_manager.activities.training.mlflow')
|
|
def test_train_model_without_logger_does_not_set_wrapper_logger(_mock_mlflow, training):
|
|
"""Covers branch where activity logger is None."""
|
|
training.logger = None
|
|
tp = TrainModelParams.from_dict(
|
|
{
|
|
**_minimal_params_dict(),
|
|
'model_metadata': {'schemas': {'components': {'schemas': {}}}},
|
|
}
|
|
)
|
|
train_df = pd.DataFrame({'a': [1.0, 2.0], 't': [1.0, 2.0]})
|
|
val_df = pd.DataFrame({'a': [1.0], 't': [1.0]})
|
|
tmr = TrainModelResult(params=tp, train_data=train_df, val_data=val_df)
|
|
|
|
training.minio_repository.download_file = MagicMock(return_value=b'csv')
|
|
training.data_manager_repository.prepare_training_data = MagicMock(return_value=tmr)
|
|
training.data_manager_repository.compute_regression_metrics = MagicMock(
|
|
side_effect=lambda x, _w, **_kw: x
|
|
)
|
|
|
|
def _fill_report(x, **_kw):
|
|
x.report_path = '/tmp/report.html'
|
|
x.train_data_path = '/tmp/train.csv'
|
|
x.test_data_path = '/tmp/test.csv'
|
|
x.equation_path = '/tmp/eq.json'
|
|
x.run_dir = '/tmp/run'
|
|
return x
|
|
|
|
training.data_manager_repository.generate_report = MagicMock(side_effect=_fill_report)
|
|
wrapper = MagicMock()
|
|
wrapper.transform = MagicMock(side_effect=[(train_df, None), (val_df, None)])
|
|
wrapper.predict = MagicMock(
|
|
side_effect=[(pd.DataFrame({'p': [1.0, 2.0]}), None), (pd.DataFrame({'p': [1.0]}), None)]
|
|
)
|
|
wrapper.store_model = MagicMock()
|
|
training.plugin_store.get_model = MagicMock(return_value=wrapper)
|
|
|
|
@contextmanager
|
|
def _run_ctx(*_a, **_k):
|
|
info = MagicMock()
|
|
info.run_name = 'n'
|
|
info.run_id = 'i'
|
|
yield info
|
|
|
|
training.mlflow_repository.start_run = _run_ctx
|
|
|
|
training.train_model({'metadata': {}, 'train_params': tp.to_dict()})
|
|
|
|
|
|
@patch('model_manager.activities.training.mlflow')
|
|
def test_train_model_train_params_as_dict(mock_mlflow, training):
|
|
"""train_params may arrive as dict and is coerced via TrainModelParams.from_dict."""
|
|
d = {
|
|
**_minimal_params_dict(),
|
|
'model_metadata': {'schemas': {'components': {'schemas': {}}}},
|
|
}
|
|
train_df = pd.DataFrame({'a': [1.0, 2.0], 't': [1.0, 2.0]})
|
|
val_df = pd.DataFrame({'a': [1.0], 't': [1.0]})
|
|
tp = TrainModelParams.from_dict(d)
|
|
tmr = TrainModelResult(params=tp, train_data=train_df, val_data=val_df)
|
|
|
|
training.minio_repository.download_file = MagicMock(return_value=b'csv')
|
|
training.data_manager_repository.prepare_training_data = MagicMock(return_value=tmr)
|
|
training.data_manager_repository.compute_regression_metrics = MagicMock(
|
|
side_effect=lambda x, _w, **_kw: x
|
|
)
|
|
|
|
def _fill_report2(x, **_kw):
|
|
x.report_path = '/tmp/report.html'
|
|
x.train_data_path = '/tmp/train.csv'
|
|
x.test_data_path = '/tmp/test.csv'
|
|
x.run_dir = '/tmp/run'
|
|
return x
|
|
|
|
training.data_manager_repository.generate_report = MagicMock(side_effect=_fill_report2)
|
|
wrapper = MagicMock()
|
|
wrapper.transform = MagicMock(side_effect=[(train_df, None), (val_df, None)])
|
|
wrapper.predict = MagicMock(
|
|
side_effect=[(pd.DataFrame({'p': [1.0, 2.0]}), None), (pd.DataFrame({'p': [1.0]}), None)]
|
|
)
|
|
wrapper.store_model = MagicMock()
|
|
training.plugin_store.get_model = MagicMock(return_value=wrapper)
|
|
|
|
@contextmanager
|
|
def _run_ctx(*_a, **_k):
|
|
info = MagicMock()
|
|
info.run_name = 'n'
|
|
info.run_id = 'i'
|
|
yield info
|
|
|
|
training.mlflow_repository.start_run = _run_ctx
|
|
|
|
training.train_model({'metadata': {}, 'train_params': d})
|
|
mock_mlflow.log_artifact.assert_called()
|
|
|
|
|
|
@patch('model_manager.activities.training.mlflow')
|
|
def test_train_model_downloads_validation_file_when_set(mock_mlflow, training):
|
|
"""Second MinIO download when val_file_name is set (covers val_bytes branch)."""
|
|
d = {
|
|
**_minimal_params_dict(),
|
|
'val_file_name': 'val.csv',
|
|
'model_metadata': {'schemas': {'components': {'schemas': {}}}},
|
|
}
|
|
tp = TrainModelParams.from_dict(d)
|
|
train_df = pd.DataFrame({'a': [1.0, 2.0], 't': [1.0, 2.0]})
|
|
val_df = pd.DataFrame({'a': [1.0], 't': [1.0]})
|
|
tmr = TrainModelResult(params=tp, train_data=train_df, val_data=val_df)
|
|
|
|
def _dl(object_name, **_kwargs):
|
|
if object_name == tp.file_name:
|
|
return b'train'
|
|
if object_name == 'val.csv':
|
|
return b'val'
|
|
raise AssertionError(object_name)
|
|
|
|
training.minio_repository.download_file = MagicMock(side_effect=_dl)
|
|
training.data_manager_repository.prepare_training_data = MagicMock(return_value=tmr)
|
|
training.data_manager_repository.compute_regression_metrics = MagicMock(
|
|
side_effect=lambda x, _w, **_kw: x
|
|
)
|
|
|
|
def _fill(x, **_kw):
|
|
x.report_path = '/tmp/report.html'
|
|
x.train_data_path = '/tmp/train.csv'
|
|
x.test_data_path = '/tmp/test.csv'
|
|
x.run_dir = '/tmp/run'
|
|
return x
|
|
|
|
training.data_manager_repository.generate_report = MagicMock(side_effect=_fill)
|
|
wrapper = MagicMock()
|
|
wrapper.transform = MagicMock(side_effect=[(train_df, None), (val_df, None)])
|
|
wrapper.predict = MagicMock(
|
|
side_effect=[(pd.DataFrame({'p': [1.0, 2.0]}), None), (pd.DataFrame({'p': [1.0]}), None)]
|
|
)
|
|
wrapper.store_model = MagicMock()
|
|
training.plugin_store.get_model = MagicMock(return_value=wrapper)
|
|
|
|
@contextmanager
|
|
def _run_ctx(*_a, **_k):
|
|
info = MagicMock()
|
|
info.run_name = 'n'
|
|
info.run_id = 'i'
|
|
yield info
|
|
|
|
training.mlflow_repository.start_run = _run_ctx
|
|
|
|
training.train_model({'metadata': {}, 'train_params': tp.to_dict()})
|
|
assert training.minio_repository.download_file.call_count == 2
|
|
mock_mlflow.log_artifact.assert_called()
|
|
|
|
|
|
def test_train_model_value_error_when_paths_missing_after_report(training):
|
|
"""Raises ValueError when report paths are not populated after generate_report."""
|
|
tp = TrainModelParams.from_dict(
|
|
{
|
|
**_minimal_params_dict(),
|
|
'model_metadata': {'schemas': {'components': {'schemas': {}}}},
|
|
}
|
|
)
|
|
train_df = pd.DataFrame({'a': [1.0, 2.0], 't': [1.0, 2.0]})
|
|
val_df = pd.DataFrame({'a': [1.0], 't': [1.0]})
|
|
tmr = TrainModelResult(params=tp, train_data=train_df, val_data=val_df)
|
|
|
|
training.minio_repository.download_file = MagicMock(return_value=b'x')
|
|
training.data_manager_repository.prepare_training_data = MagicMock(return_value=tmr)
|
|
training.data_manager_repository.compute_regression_metrics = MagicMock(
|
|
side_effect=lambda x, _w, **_kw: x
|
|
)
|
|
training.data_manager_repository.generate_report = MagicMock(return_value=tmr)
|
|
wrapper = MagicMock()
|
|
wrapper.transform = MagicMock(side_effect=[(train_df, None), (val_df, None)])
|
|
wrapper.predict = MagicMock(
|
|
side_effect=[(pd.DataFrame({'p': [1.0, 2.0]}), None), (pd.DataFrame({'p': [1.0]}), None)]
|
|
)
|
|
training.plugin_store.get_model = MagicMock(return_value=wrapper)
|
|
|
|
@contextmanager
|
|
def _run_ctx(*_a, **_k):
|
|
info = MagicMock()
|
|
info.run_name = 'n'
|
|
info.run_id = 'i'
|
|
yield info
|
|
|
|
training.mlflow_repository.start_run = _run_ctx
|
|
training.send_notification = MagicMock()
|
|
|
|
with pytest.raises(ValueError, match='Report path'):
|
|
training.train_model({'metadata': {}, 'train_params': tp.to_dict()})
|