test(simple_metrics): rewrite tests against RegressionMetrics dispatch
Rewrites 8 existing tests to mock RegressionMetrics instead of verifying manual numpy math. Adds 2 new tests: - r2 excluded + warning when model_type is non-linear - r2 included when model_type is absent (backward compat) Replaces silent-discard test with ValueError propagation test. SIENTIAPDE-1986
This commit is contained in:
@@ -45,6 +45,23 @@ metadata = {
|
|||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _mock_regression_metrics_class(calculate_return):
|
||||||
|
"""Returns a patch context manager that mocks RegressionMetrics."""
|
||||||
|
mock_instance = MagicMock()
|
||||||
|
mock_instance.calculate.return_value = calculate_return
|
||||||
|
mock_class = MagicMock(return_value=mock_instance)
|
||||||
|
mock_class.is_r2_supported = MagicMock(return_value=True)
|
||||||
|
mock_class.supported_metrics = MagicMock(return_value=['rmse', 'mse', 'mae', 'r2'])
|
||||||
|
return (
|
||||||
|
patch(
|
||||||
|
'laborious.activities.model_metrics.RegressionMetrics',
|
||||||
|
mock_class,
|
||||||
|
),
|
||||||
|
mock_class,
|
||||||
|
mock_instance,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def test_calculate_drift_invalid_chunk_period(model_metrics_activity):
|
def test_calculate_drift_invalid_chunk_period(model_metrics_activity):
|
||||||
# Arrange
|
# Arrange
|
||||||
input_data = {
|
input_data = {
|
||||||
@@ -693,7 +710,6 @@ def test_get_drift_metrics_dataframe_error(
|
|||||||
|
|
||||||
|
|
||||||
def test_calculate_simple_metrics_success_all_metrics(model_metrics_activity):
|
def test_calculate_simple_metrics_success_all_metrics(model_metrics_activity):
|
||||||
# Arrange
|
|
||||||
target_data = DataFrame(
|
target_data = DataFrame(
|
||||||
{
|
{
|
||||||
'timestamp': ['2023-05-26 11:12:27', '2023-05-26 11:12:28', '2023-05-26 11:12:29'],
|
'timestamp': ['2023-05-26 11:12:27', '2023-05-26 11:12:28', '2023-05-26 11:12:29'],
|
||||||
@@ -710,28 +726,28 @@ def test_calculate_simple_metrics_success_all_metrics(model_metrics_activity):
|
|||||||
'interval_minutes': 5,
|
'interval_minutes': 5,
|
||||||
}
|
}
|
||||||
|
|
||||||
# Act
|
patcher, mock_class, mock_instance = _mock_regression_metrics_class(
|
||||||
result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data))
|
[
|
||||||
|
{'metric': 'rmse', 'value': 0.1},
|
||||||
# Assert
|
{'metric': 'mse', 'value': 0.01},
|
||||||
assert len(result['metric']) == 4
|
{'metric': 'mae', 'value': 0.1},
|
||||||
assert 'rmse' in result['metric'].values
|
{'metric': 'r2', 'value': 0.99},
|
||||||
assert 'mse' in result['metric'].values
|
]
|
||||||
assert 'mae' in result['metric'].values
|
|
||||||
assert 'r2' in result['metric'].values
|
|
||||||
assert all(model_id == 'test_model_id' for model_id in result['model_id'].values)
|
|
||||||
assert all(timestamp == '2023-05-26 11:12:29' for timestamp in result['timestamp'].values)
|
|
||||||
assert all(data_size == 3 for data_size in result['data_size'].values)
|
|
||||||
assert all(interval_minutes == 5 for interval_minutes in result['interval_minutes'].values)
|
|
||||||
model_metrics_activity.info.assert_called_once_with(
|
|
||||||
"Calculating simple metrics for model test_model_id: ['rmse', 'mse', 'mae', 'r2']",
|
|
||||||
metadata['metadata'],
|
|
||||||
)
|
)
|
||||||
model_metrics_activity.debug.assert_called_once()
|
|
||||||
|
with patcher:
|
||||||
|
result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data))
|
||||||
|
|
||||||
|
assert len(result) == 4
|
||||||
|
assert set(result['metric'].values) == {'rmse', 'mse', 'mae', 'r2'}
|
||||||
|
assert all(mid == 'test_model_id' for mid in result['model_id'].values)
|
||||||
|
assert all(ts == '2023-05-26 11:12:29' for ts in result['timestamp'].values)
|
||||||
|
assert all(ds == 3 for ds in result['data_size'].values)
|
||||||
|
assert all(im == 5 for im in result['interval_minutes'].values)
|
||||||
|
mock_instance.calculate.assert_called_once_with(['rmse', 'mse', 'mae', 'r2'])
|
||||||
|
|
||||||
|
|
||||||
def test_calculate_simple_metrics_success_rmse_only(model_metrics_activity):
|
def test_calculate_simple_metrics_success_rmse_only(model_metrics_activity):
|
||||||
# Arrange
|
|
||||||
target_data = DataFrame(
|
target_data = DataFrame(
|
||||||
{
|
{
|
||||||
'timestamp': ['2023-05-26 11:12:27', '2023-05-26 11:12:28'],
|
'timestamp': ['2023-05-26 11:12:27', '2023-05-26 11:12:28'],
|
||||||
@@ -748,23 +764,23 @@ def test_calculate_simple_metrics_success_rmse_only(model_metrics_activity):
|
|||||||
'interval_minutes': 5,
|
'interval_minutes': 5,
|
||||||
}
|
}
|
||||||
|
|
||||||
# Act
|
patcher, mock_class, mock_instance = _mock_regression_metrics_class(
|
||||||
result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data))
|
[{'metric': 'rmse', 'value': 0.1}]
|
||||||
|
)
|
||||||
|
|
||||||
# Assert
|
with patcher:
|
||||||
assert len(result['metric']) == 1
|
result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data))
|
||||||
|
|
||||||
|
assert len(result) == 1
|
||||||
assert result['metric'].values[0] == 'rmse'
|
assert result['metric'].values[0] == 'rmse'
|
||||||
assert result['model_id'].values[0] == 'test_model_id'
|
assert result['model_id'].values[0] == 'test_model_id'
|
||||||
assert result['timestamp'].values[0] == '2023-05-26 11:12:28'
|
assert result['timestamp'].values[0] == '2023-05-26 11:12:28'
|
||||||
assert result['data_size'].values[0] == 2
|
assert result['data_size'].values[0] == 2
|
||||||
assert result['interval_minutes'].values[0] == 5
|
assert result['interval_minutes'].values[0] == 5
|
||||||
model_metrics_activity.info.assert_called_once_with(
|
mock_instance.calculate.assert_called_once_with(['rmse'])
|
||||||
"Calculating simple metrics for model test_model_id: ['rmse']", metadata['metadata']
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def test_calculate_simple_metrics_success_mse_only(model_metrics_activity):
|
def test_calculate_simple_metrics_success_mse_only(model_metrics_activity):
|
||||||
# Arrange
|
|
||||||
target_data = DataFrame(
|
target_data = DataFrame(
|
||||||
{
|
{
|
||||||
'timestamp': ['2023-05-26 11:12:27', '2023-05-26 11:12:28'],
|
'timestamp': ['2023-05-26 11:12:27', '2023-05-26 11:12:28'],
|
||||||
@@ -781,23 +797,23 @@ def test_calculate_simple_metrics_success_mse_only(model_metrics_activity):
|
|||||||
'interval_minutes': 5,
|
'interval_minutes': 5,
|
||||||
}
|
}
|
||||||
|
|
||||||
# Act
|
patcher, mock_class, mock_instance = _mock_regression_metrics_class(
|
||||||
result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data))
|
[{'metric': 'mse', 'value': 0.01}]
|
||||||
|
)
|
||||||
|
|
||||||
# Assert
|
with patcher:
|
||||||
assert len(result['metric']) == 1
|
result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data))
|
||||||
|
|
||||||
|
assert len(result) == 1
|
||||||
assert result['metric'].values[0] == 'mse'
|
assert result['metric'].values[0] == 'mse'
|
||||||
assert result['model_id'].values[0] == 'test_model_id'
|
assert result['model_id'].values[0] == 'test_model_id'
|
||||||
assert result['timestamp'].values[0] == '2023-05-26 11:12:28'
|
assert result['timestamp'].values[0] == '2023-05-26 11:12:28'
|
||||||
assert result['data_size'].values[0] == 2
|
assert result['data_size'].values[0] == 2
|
||||||
assert result['interval_minutes'].values[0] == 5
|
assert result['interval_minutes'].values[0] == 5
|
||||||
model_metrics_activity.info.assert_called_once_with(
|
mock_instance.calculate.assert_called_once_with(['mse'])
|
||||||
"Calculating simple metrics for model test_model_id: ['mse']", metadata['metadata']
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def test_calculate_simple_metrics_success_mae_only(model_metrics_activity):
|
def test_calculate_simple_metrics_success_mae_only(model_metrics_activity):
|
||||||
# Arrange
|
|
||||||
target_data = DataFrame(
|
target_data = DataFrame(
|
||||||
{
|
{
|
||||||
'timestamp': ['2023-05-26 11:12:27', '2023-05-26 11:12:28'],
|
'timestamp': ['2023-05-26 11:12:27', '2023-05-26 11:12:28'],
|
||||||
@@ -814,23 +830,23 @@ def test_calculate_simple_metrics_success_mae_only(model_metrics_activity):
|
|||||||
'interval_minutes': 5,
|
'interval_minutes': 5,
|
||||||
}
|
}
|
||||||
|
|
||||||
# Act
|
patcher, mock_class, mock_instance = _mock_regression_metrics_class(
|
||||||
result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data))
|
[{'metric': 'mae', 'value': 0.1}]
|
||||||
|
)
|
||||||
|
|
||||||
# Assert
|
with patcher:
|
||||||
assert len(result['metric']) == 1
|
result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data))
|
||||||
|
|
||||||
|
assert len(result) == 1
|
||||||
assert result['metric'].values[0] == 'mae'
|
assert result['metric'].values[0] == 'mae'
|
||||||
assert result['model_id'].values[0] == 'test_model_id'
|
assert result['model_id'].values[0] == 'test_model_id'
|
||||||
assert result['timestamp'].values[0] == '2023-05-26 11:12:28'
|
assert result['timestamp'].values[0] == '2023-05-26 11:12:28'
|
||||||
assert result['data_size'].values[0] == 2
|
assert result['data_size'].values[0] == 2
|
||||||
assert result['interval_minutes'].values[0] == 5
|
assert result['interval_minutes'].values[0] == 5
|
||||||
model_metrics_activity.info.assert_called_once_with(
|
mock_instance.calculate.assert_called_once_with(['mae'])
|
||||||
"Calculating simple metrics for model test_model_id: ['mae']", metadata['metadata']
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def test_calculate_simple_metrics_success_r2_only(model_metrics_activity):
|
def test_calculate_simple_metrics_success_r2_only(model_metrics_activity):
|
||||||
# Arrange
|
|
||||||
target_data = DataFrame(
|
target_data = DataFrame(
|
||||||
{
|
{
|
||||||
'timestamp': ['2023-05-26 11:12:27', '2023-05-26 11:12:28'],
|
'timestamp': ['2023-05-26 11:12:27', '2023-05-26 11:12:28'],
|
||||||
@@ -847,24 +863,24 @@ def test_calculate_simple_metrics_success_r2_only(model_metrics_activity):
|
|||||||
'interval_minutes': 5,
|
'interval_minutes': 5,
|
||||||
}
|
}
|
||||||
|
|
||||||
# Act
|
patcher, mock_class, mock_instance = _mock_regression_metrics_class(
|
||||||
result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data))
|
[{'metric': 'r2', 'value': 0.95}]
|
||||||
|
)
|
||||||
|
|
||||||
# Assert
|
with patcher:
|
||||||
assert len(result['metric']) == 1
|
result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data))
|
||||||
|
|
||||||
|
assert len(result) == 1
|
||||||
assert result['metric'].values[0] == 'r2'
|
assert result['metric'].values[0] == 'r2'
|
||||||
assert result['model_id'].values[0] == 'test_model_id'
|
assert result['model_id'].values[0] == 'test_model_id'
|
||||||
assert result['timestamp'].values[0] == '2023-05-26 11:12:28'
|
assert result['timestamp'].values[0] == '2023-05-26 11:12:28'
|
||||||
assert result['data_size'].values[0] == 2
|
assert result['data_size'].values[0] == 2
|
||||||
assert result['interval_minutes'].values[0] == 5
|
assert result['interval_minutes'].values[0] == 5
|
||||||
model_metrics_activity.info.assert_called_once_with(
|
mock_instance.calculate.assert_called_once_with(['r2'])
|
||||||
"Calculating simple metrics for model test_model_id: ['r2']", metadata['metadata']
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def test_calculate_simple_metrics_r2_zero_ss_tot(model_metrics_activity):
|
def test_calculate_simple_metrics_r2_zero_variance_delegates_to_lib(model_metrics_activity):
|
||||||
# Arrange
|
"""r2 zero-variance is now the lib's responsibility. Activity just passes through."""
|
||||||
# All target values are the same, so ss_tot will be 0
|
|
||||||
target_data = DataFrame(
|
target_data = DataFrame(
|
||||||
{
|
{
|
||||||
'timestamp': ['2023-05-26 11:12:27', '2023-05-26 11:12:28'],
|
'timestamp': ['2023-05-26 11:12:27', '2023-05-26 11:12:28'],
|
||||||
@@ -881,24 +897,18 @@ def test_calculate_simple_metrics_r2_zero_ss_tot(model_metrics_activity):
|
|||||||
'interval_minutes': 5,
|
'interval_minutes': 5,
|
||||||
}
|
}
|
||||||
|
|
||||||
# Act
|
patcher, mock_class, mock_instance = _mock_regression_metrics_class(
|
||||||
result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data))
|
[{'metric': 'r2', 'value': 0.0}]
|
||||||
|
|
||||||
# Assert
|
|
||||||
assert len(result['metric']) == 1
|
|
||||||
assert result['metric'].values[0] == 'r2'
|
|
||||||
assert result['value'].values[0] == 0.0 # Should return 0.0 when ss_tot == 0
|
|
||||||
assert result['model_id'].values[0] == 'test_model_id'
|
|
||||||
assert result['timestamp'].values[0] == '2023-05-26 11:12:28'
|
|
||||||
assert result['data_size'].values[0] == 2
|
|
||||||
assert result['interval_minutes'].values[0] == 5
|
|
||||||
model_metrics_activity.info.assert_called_once_with(
|
|
||||||
"Calculating simple metrics for model test_model_id: ['r2']", metadata['metadata']
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
with patcher:
|
||||||
|
result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data))
|
||||||
|
|
||||||
|
assert result['value'].values[0] == 0.0
|
||||||
|
mock_instance.calculate.assert_called_once_with(['r2'])
|
||||||
|
|
||||||
|
|
||||||
def test_calculate_simple_metrics_success_multiple_metrics_subset(model_metrics_activity):
|
def test_calculate_simple_metrics_success_multiple_metrics_subset(model_metrics_activity):
|
||||||
# Arrange
|
|
||||||
target_data = DataFrame(
|
target_data = DataFrame(
|
||||||
{
|
{
|
||||||
'timestamp': ['2023-05-26 11:12:27', '2023-05-26 11:12:28', '2023-05-26 11:12:29'],
|
'timestamp': ['2023-05-26 11:12:27', '2023-05-26 11:12:28', '2023-05-26 11:12:29'],
|
||||||
@@ -915,23 +925,26 @@ def test_calculate_simple_metrics_success_multiple_metrics_subset(model_metrics_
|
|||||||
'interval_minutes': 5,
|
'interval_minutes': 5,
|
||||||
}
|
}
|
||||||
|
|
||||||
# Act
|
patcher, mock_class, mock_instance = _mock_regression_metrics_class(
|
||||||
result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data))
|
[
|
||||||
|
{'metric': 'rmse', 'value': 0.1},
|
||||||
# Assert
|
{'metric': 'mae', 'value': 0.1},
|
||||||
assert len(result['metric']) == 2
|
]
|
||||||
assert 'rmse' in result['metric'].values
|
|
||||||
assert 'mae' in result['metric'].values
|
|
||||||
assert all(model_id == 'test_model_id' for model_id in result['model_id'].values)
|
|
||||||
assert all(timestamp == '2023-05-26 11:12:29' for timestamp in result['timestamp'].values)
|
|
||||||
assert all(data_size == 3 for data_size in result['data_size'].values)
|
|
||||||
assert all(interval_minutes == 5 for interval_minutes in result['interval_minutes'].values)
|
|
||||||
model_metrics_activity.info.assert_called_once_with(
|
|
||||||
"Calculating simple metrics for model test_model_id: ['rmse', 'mae']", metadata['metadata']
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
with patcher:
|
||||||
|
result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data))
|
||||||
|
|
||||||
def test_calculate_simple_metrics_unknown_metric_ignored(model_metrics_activity):
|
assert len(result) == 2
|
||||||
|
assert set(result['metric'].values) == {'rmse', 'mae'}
|
||||||
|
assert all(mid == 'test_model_id' for mid in result['model_id'].values)
|
||||||
|
assert all(ts == '2023-05-26 11:12:29' for ts in result['timestamp'].values)
|
||||||
|
assert all(ds == 3 for ds in result['data_size'].values)
|
||||||
|
assert all(im == 5 for im in result['interval_minutes'].values)
|
||||||
|
mock_instance.calculate.assert_called_once_with(['rmse', 'mae'])
|
||||||
|
|
||||||
|
|
||||||
|
def test_calculate_simple_metrics_unknown_metric_raises(model_metrics_activity):
|
||||||
target_data = DataFrame(
|
target_data = DataFrame(
|
||||||
{
|
{
|
||||||
'timestamp': ['2023-05-26 11:12:27', '2023-05-26 11:12:28'],
|
'timestamp': ['2023-05-26 11:12:27', '2023-05-26 11:12:28'],
|
||||||
@@ -948,7 +961,79 @@ def test_calculate_simple_metrics_unknown_metric_ignored(model_metrics_activity)
|
|||||||
'interval_minutes': 5,
|
'interval_minutes': 5,
|
||||||
}
|
}
|
||||||
|
|
||||||
result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data))
|
mock_instance = MagicMock()
|
||||||
|
mock_instance.calculate.side_effect = ValueError('Unknown metric: unknown_metric')
|
||||||
|
mock_class = MagicMock(return_value=mock_instance)
|
||||||
|
mock_class.is_r2_supported = MagicMock(return_value=True)
|
||||||
|
|
||||||
assert len(result['metric']) == 1
|
with patch('laborious.activities.model_metrics.RegressionMetrics', mock_class):
|
||||||
|
with raises(ValueError, match='Unknown metric'):
|
||||||
|
model_metrics_activity.calculate_simple_metrics(input_data)
|
||||||
|
|
||||||
|
|
||||||
|
def test_calculate_simple_metrics_r2_skipped_for_nonlinear_model(model_metrics_activity):
|
||||||
|
"""When model_type is non-linear, r2 is excluded and a warning notification fires."""
|
||||||
|
target_data = DataFrame(
|
||||||
|
{
|
||||||
|
'timestamp': ['2023-05-26 11:12:27', '2023-05-26 11:12:28'],
|
||||||
|
'target': [1.0, 2.0],
|
||||||
|
'prediction': [1.1, 2.1],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
input_data = {
|
||||||
|
**metadata,
|
||||||
|
'model_id': 'test_model_id',
|
||||||
|
'target_data': target_data.to_dict(),
|
||||||
|
'metrics': ['rmse', 'r2'],
|
||||||
|
'interval_minutes': 5,
|
||||||
|
'model_type': 'XGBoost',
|
||||||
|
}
|
||||||
|
|
||||||
|
mock_instance = MagicMock()
|
||||||
|
mock_instance.calculate.return_value = [{'metric': 'rmse', 'value': 0.1}]
|
||||||
|
mock_class = MagicMock(return_value=mock_instance)
|
||||||
|
mock_class.is_r2_supported = MagicMock(return_value=False)
|
||||||
|
|
||||||
|
with patch('laborious.activities.model_metrics.RegressionMetrics', mock_class):
|
||||||
|
result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data))
|
||||||
|
|
||||||
|
assert len(result) == 1
|
||||||
assert result['metric'].values[0] == 'rmse'
|
assert result['metric'].values[0] == 'rmse'
|
||||||
|
mock_class.is_r2_supported.assert_called_once_with('XGBoost')
|
||||||
|
mock_instance.calculate.assert_called_once_with(['rmse'])
|
||||||
|
model_metrics_activity.warning.assert_called_once()
|
||||||
|
model_metrics_activity.send_notification.assert_called_once()
|
||||||
|
call_kwargs = model_metrics_activity.send_notification.call_args.kwargs
|
||||||
|
assert call_kwargs['notification_id'] == 'SIMPLE_METRICS_R2_UNSUPPORTED'
|
||||||
|
assert call_kwargs['level'] == NotificationLevel.WARNING
|
||||||
|
|
||||||
|
|
||||||
|
def test_calculate_simple_metrics_no_model_type_includes_r2(model_metrics_activity):
|
||||||
|
"""When model_type is None (legacy input), r2 is included without check."""
|
||||||
|
target_data = DataFrame(
|
||||||
|
{
|
||||||
|
'timestamp': ['2023-05-26 11:12:27', '2023-05-26 11:12:28'],
|
||||||
|
'target': [1.0, 2.0],
|
||||||
|
'prediction': [1.1, 2.1],
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
input_data = {
|
||||||
|
**metadata,
|
||||||
|
'model_id': 'test_model_id',
|
||||||
|
'target_data': target_data.to_dict(),
|
||||||
|
'metrics': ['r2'],
|
||||||
|
'interval_minutes': 5,
|
||||||
|
# no model_type key
|
||||||
|
}
|
||||||
|
|
||||||
|
patcher, mock_class, mock_instance = _mock_regression_metrics_class(
|
||||||
|
[{'metric': 'r2', 'value': 0.95}]
|
||||||
|
)
|
||||||
|
|
||||||
|
with patcher:
|
||||||
|
result = DataFrame(model_metrics_activity.calculate_simple_metrics(input_data))
|
||||||
|
|
||||||
|
assert result['metric'].values[0] == 'r2'
|
||||||
|
mock_class.is_r2_supported.assert_not_called()
|
||||||
|
|||||||
Reference in New Issue
Block a user