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.
This commit is contained in:
@@ -1,7 +1,7 @@
|
||||
"""Unit tests for ExperimentTracking class with 100% coverage."""
|
||||
|
||||
import asyncio
|
||||
from unittest.mock import MagicMock, patch
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
@@ -22,12 +22,8 @@ def mock_notification_handler():
|
||||
def mock_metrics_controller():
|
||||
"""Create a mock metrics controller."""
|
||||
controller = MagicMock()
|
||||
|
||||
# Make shutdown an async coroutine
|
||||
async def mock_shutdown():
|
||||
pass
|
||||
|
||||
controller.shutdown = mock_shutdown
|
||||
controller.shutdown = AsyncMock()
|
||||
controller.emit = AsyncMock()
|
||||
return controller
|
||||
|
||||
|
||||
@@ -259,7 +255,7 @@ def test_update_experiment_run_status_missing_status(
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
|
||||
et.send_notification = MagicMock()
|
||||
et.send_notification_async = AsyncMock()
|
||||
|
||||
input_data = {
|
||||
'metadata': {'workflow_id': 'test-123'},
|
||||
@@ -270,7 +266,7 @@ def test_update_experiment_run_status_missing_status(
|
||||
with pytest.raises(RuntimeError):
|
||||
asyncio.run(et.update_experiment_run(input_data))
|
||||
|
||||
et.send_notification.assert_called_once()
|
||||
et.send_notification_async.assert_awaited_once()
|
||||
|
||||
|
||||
def test_update_experiment_run_status_with_error_success(
|
||||
@@ -381,7 +377,7 @@ def test_update_experiment_run_status_with_error_missing_error_message(
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
|
||||
et.send_notification = MagicMock()
|
||||
et.send_notification_async = AsyncMock()
|
||||
|
||||
input_data = {
|
||||
'metadata': {'workflow_id': 'test-123'},
|
||||
@@ -393,7 +389,7 @@ def test_update_experiment_run_status_with_error_missing_error_message(
|
||||
with pytest.raises(RuntimeError):
|
||||
asyncio.run(et.update_experiment_run(input_data))
|
||||
|
||||
et.send_notification.assert_called_once()
|
||||
et.send_notification_async.assert_awaited_once()
|
||||
|
||||
|
||||
def test_update_experiment_run_model_saved_success(
|
||||
@@ -461,7 +457,7 @@ def test_update_experiment_run_model_saved_missing_run_name(
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
|
||||
et.send_notification = MagicMock()
|
||||
et.send_notification_async = AsyncMock()
|
||||
|
||||
input_data = {
|
||||
'metadata': {'workflow_id': 'test-123'},
|
||||
@@ -473,7 +469,7 @@ def test_update_experiment_run_model_saved_missing_run_name(
|
||||
with pytest.raises(RuntimeError):
|
||||
asyncio.run(et.update_experiment_run(input_data))
|
||||
|
||||
et.send_notification.assert_called_once()
|
||||
et.send_notification_async.assert_awaited_once()
|
||||
|
||||
|
||||
def test_update_experiment_run_invalid_update_type(
|
||||
@@ -495,7 +491,7 @@ def test_update_experiment_run_invalid_update_type(
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
|
||||
et.send_notification = MagicMock()
|
||||
et.send_notification_async = AsyncMock()
|
||||
|
||||
input_data = {
|
||||
'metadata': {'workflow_id': 'test-123'},
|
||||
@@ -506,7 +502,7 @@ def test_update_experiment_run_invalid_update_type(
|
||||
with pytest.raises(RuntimeError):
|
||||
asyncio.run(et.update_experiment_run(input_data))
|
||||
|
||||
et.send_notification.assert_called_once()
|
||||
et.send_notification_async.assert_awaited_once()
|
||||
|
||||
|
||||
def test_update_experiment_run_no_rows_updated(
|
||||
@@ -532,7 +528,7 @@ def test_update_experiment_run_no_rows_updated(
|
||||
return {'rowcount': 0}
|
||||
|
||||
et._execute_update = mock_execute_update
|
||||
et.send_notification = MagicMock()
|
||||
et.send_notification_async = AsyncMock()
|
||||
|
||||
input_data = {
|
||||
'metadata': {'workflow_id': 'test-123'},
|
||||
@@ -544,7 +540,7 @@ def test_update_experiment_run_no_rows_updated(
|
||||
with pytest.raises(RuntimeError):
|
||||
asyncio.run(et.update_experiment_run(input_data))
|
||||
|
||||
et.send_notification.assert_called_once()
|
||||
et.send_notification_async.assert_awaited_once()
|
||||
|
||||
|
||||
def test_update_experiment_run_status_with_error_missing_status(
|
||||
@@ -566,7 +562,7 @@ def test_update_experiment_run_status_with_error_missing_status(
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
|
||||
et.send_notification = MagicMock()
|
||||
et.send_notification_async = AsyncMock()
|
||||
|
||||
input_data = {
|
||||
'metadata': {'workflow_id': 'test-123'},
|
||||
@@ -578,7 +574,7 @@ def test_update_experiment_run_status_with_error_missing_status(
|
||||
with pytest.raises(RuntimeError):
|
||||
asyncio.run(et.update_experiment_run(input_data))
|
||||
|
||||
et.send_notification.assert_called_once()
|
||||
et.send_notification_async.assert_awaited_once()
|
||||
|
||||
|
||||
def test_update_experiment_run_model_saved_missing_status(
|
||||
@@ -600,7 +596,7 @@ def test_update_experiment_run_model_saved_missing_status(
|
||||
metrics_controller=mock_metrics_controller,
|
||||
)
|
||||
|
||||
et.send_notification = MagicMock()
|
||||
et.send_notification_async = AsyncMock()
|
||||
|
||||
input_data = {
|
||||
'metadata': {'workflow_id': 'test-123'},
|
||||
@@ -612,7 +608,7 @@ def test_update_experiment_run_model_saved_missing_status(
|
||||
with pytest.raises(RuntimeError):
|
||||
asyncio.run(et.update_experiment_run(input_data))
|
||||
|
||||
et.send_notification.assert_called_once()
|
||||
et.send_notification_async.assert_awaited_once()
|
||||
|
||||
|
||||
def test_experiment_tracking_del_with_engine_no_super_del(
|
||||
|
||||
Reference in New Issue
Block a user