feat: log regression metrics as parameters in Training class
- Added a method to persist computed regression metrics (MSE, MAE, R²) as MLflow parameters during model training, enhancing model evaluation and tracking. - Updated the Training class to log the equation path if available, improving artifact management.
This commit is contained in:
@@ -132,7 +132,12 @@ def test_train_model_success_serializes_result(mock_mlflow, training):
|
||||
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: setattr(x, 'mse_val', 0.1) or x
|
||||
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):
|
||||
@@ -145,12 +150,7 @@ def test_train_model_success_serializes_result(mock_mlflow, training):
|
||||
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.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)])
|
||||
@@ -167,12 +167,64 @@ def test_train_model_success_serializes_result(mock_mlflow, training):
|
||||
training.mlflow_repository.start_run = _run_ctx
|
||||
|
||||
out = training.train_model({'metadata': {'pod': 'p'}, 'train_params': tp.to_dict()})
|
||||
assert out['run_name'] == 'run-n'
|
||||
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."""
|
||||
@@ -188,7 +240,7 @@ def test_train_model_train_params_as_dict(mock_mlflow, training):
|
||||
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: x
|
||||
side_effect=lambda x, _w, **_kw: x
|
||||
)
|
||||
|
||||
def _fill_report2(x, **_kw):
|
||||
@@ -243,7 +295,7 @@ def test_train_model_downloads_validation_file_when_set(mock_mlflow, training):
|
||||
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: x
|
||||
side_effect=lambda x, _w, **_kw: x
|
||||
)
|
||||
|
||||
def _fill(x, **_kw):
|
||||
@@ -291,7 +343,7 @@ def test_train_model_value_error_when_paths_missing_after_report(training):
|
||||
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: x
|
||||
side_effect=lambda x, _w, **_kw: x
|
||||
)
|
||||
training.data_manager_repository.generate_report = MagicMock(return_value=tmr)
|
||||
wrapper = MagicMock()
|
||||
|
||||
Reference in New Issue
Block a user