Files
sientia-dataops-model-manager/tests/activities/test_training.py
vitor-aignosi 6b1df7c3a7 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.
2026-04-06 15:05:57 -03:00

316 lines
12 KiB
Python

"""Unit tests for Training activities."""
from contextlib import asynccontextmanager
from unittest.mock import AsyncMock, 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': 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': 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(),
)
@pytest.mark.asyncio
async 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 = await training.load_model_metadata(inp)
assert 'model_metadata' in out
assert out['model_metadata']['schemas']
@pytest.mark.asyncio
async def test_load_model_metadata_notifies_on_error(training):
training.plugin_store.get_model_index = MagicMock(side_effect=RuntimeError('idx'))
training.send_notification_async = AsyncMock()
inp = {**_minimal_params_dict(), 'metadata': {}}
with pytest.raises(RuntimeError, match='idx'):
await training.load_model_metadata(inp)
training.send_notification_async.assert_awaited()
@pytest.mark.asyncio
async def test_validate_train_params_success(training):
pdict = _minimal_params_dict()
pdict['model_metadata'] = {'schemas': {'components': {'schemas': {}}}}
inp = {**pdict, 'metadata': {}}
out = await training.validate_train_params(inp)
assert isinstance(out, TrainModelParams)
assert out.target_variable == 't'
@pytest.mark.asyncio
async def test_validate_train_params_notifies(training):
training.send_notification_async = AsyncMock()
inp = {'metadata': {}, 'experiment_run_id': 1}
with pytest.raises(Exception):
await training.validate_train_params(inp)
training.send_notification_async.assert_awaited()
@pytest.mark.asyncio
async 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 = AsyncMock(side_effect=OSError('minio'))
training.send_notification_async = AsyncMock()
with pytest.raises(OSError, match='minio'):
await training.train_model({'metadata': {'pod': 'x'}, 'train_params': tp})
training.send_notification_async.assert_awaited()
@pytest.mark.asyncio
async def test_cleanup_resources(training):
training.data_manager_repository.cleanup_run_directory = MagicMock()
await training.cleanup_resources({'metadata': {}, 'run_dir': '/tmp/x'})
training.data_manager_repository.cleanup_run_directory.assert_called_once_with('/tmp/x', {})
@pytest.mark.asyncio
async 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'):
await training.cleanup_resources({'metadata': {'pod': 'p'}, 'run_dir': '/tmp/x'})
training.send_notification.assert_called_once()
@pytest.mark.asyncio
@patch('model_manager.activities.training.mlflow')
async 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 = AsyncMock(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
)
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 = AsyncMock(return_value=wrapper)
@asynccontextmanager
async 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 = await training.train_model({'metadata': {'pod': 'p'}, 'train_params': tp})
assert out['run_name'] == 'run-n'
assert out['run_id'] == 'run-i'
assert out['run_dir'] == '/tmp/run'
mock_mlflow.log_artifact.assert_called()
@pytest.mark.asyncio
@patch('model_manager.activities.training.mlflow')
async 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 = AsyncMock(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)
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 = AsyncMock(return_value=wrapper)
@asynccontextmanager
async def _run_ctx(*_a, **_k):
info = MagicMock()
info.run_name = 'n'
info.run_id = 'i'
yield info
training.mlflow_repository.start_run = _run_ctx
await training.train_model({'metadata': {}, 'train_params': d})
mock_mlflow.log_artifact.assert_called()
@pytest.mark.asyncio
@patch('model_manager.activities.training.mlflow')
async 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)
async 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 = AsyncMock(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)
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 = AsyncMock(return_value=wrapper)
@asynccontextmanager
async def _run_ctx(*_a, **_k):
info = MagicMock()
info.run_name = 'n'
info.run_id = 'i'
yield info
training.mlflow_repository.start_run = _run_ctx
await training.train_model({'metadata': {}, 'train_params': tp})
assert training.minio_repository.download_file.await_count == 2
mock_mlflow.log_artifact.assert_called()
@pytest.mark.asyncio
async 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 = AsyncMock(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)
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 = AsyncMock(return_value=wrapper)
@asynccontextmanager
async 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_async = AsyncMock()
with pytest.raises(ValueError, match='Report path'):
await training.train_model({'metadata': {}, 'train_params': tp})