SIENTIAPDE-1273
Enhance data handling and export processes in Laborious workflows - Updated `gates.py` to improve data quality validation, filtering, and formatting operations, including enhanced metrics recording. - Refined `mlflow.py` to better manage model transformations and reference data retrieval from MLflow Model Registry. - Enhanced `format_and_export_prediction.py` to support separate export of transformed data, improving flexibility in data handling. - Added comprehensive test coverage for new functionalities, including transformed data formatting and retrain report generation. - Improved documentation in `README.md` to reflect changes in activities and workflows, ensuring clarity on data processing and export paths.
This commit is contained in:
@@ -546,6 +546,80 @@ async def test_format_prediction_with_timestamp_invalid_policy(gates_activity):
|
||||
raise AssertionError('Expected ValueError')
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_format_transformed_data_single_row(gates_activity):
|
||||
# Arrange
|
||||
input_data = {
|
||||
**metadata,
|
||||
'data': {
|
||||
'var1': {'2023-05-26 11:12:27': 1.0},
|
||||
'var2': {'2023-05-26 11:12:27': 2.0},
|
||||
},
|
||||
'model_id': 'test_model',
|
||||
}
|
||||
|
||||
# Act
|
||||
result = await gates_activity.format_transformed_data(input_data)
|
||||
|
||||
# Assert
|
||||
assert result['timestamp'] == {0: '2023-05-26 11:12:27', 1: '2023-05-26 11:12:27'}
|
||||
assert result['variable'] == {0: 'var1', 1: 'var2'}
|
||||
assert result['value'] == {0: 1.0, 1: 2.0}
|
||||
assert result['model_id'] == {0: 'test_model', 1: 'test_model'}
|
||||
gates_activity.info.assert_called()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_format_transformed_data_multiple_rows(gates_activity):
|
||||
# Arrange
|
||||
input_data = {
|
||||
**metadata,
|
||||
'data': {
|
||||
'var1': {
|
||||
'2023-05-26 11:12:27': 1.0,
|
||||
'2023-05-26 11:12:28': 2.0,
|
||||
},
|
||||
'var2': {
|
||||
'2023-05-26 11:12:27': 3.0,
|
||||
'2023-05-26 11:12:28': 4.0,
|
||||
},
|
||||
},
|
||||
'model_id': 'test_model',
|
||||
}
|
||||
|
||||
# Act
|
||||
result = await gates_activity.format_transformed_data(input_data)
|
||||
|
||||
# Assert
|
||||
assert len(result['timestamp']) == 4
|
||||
assert len(result['variable']) == 4
|
||||
assert len(result['value']) == 4
|
||||
assert len(result['model_id']) == 4
|
||||
assert all(v == 'test_model' for v in result['model_id'].values())
|
||||
assert set(result['variable'].values()) == {'var1', 'var2'}
|
||||
gates_activity.info.assert_called()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_format_transformed_data_empty_data(gates_activity):
|
||||
# Arrange
|
||||
input_data = {
|
||||
**metadata,
|
||||
'data': {},
|
||||
'model_id': 'test_model',
|
||||
}
|
||||
|
||||
# Act
|
||||
result = await gates_activity.format_transformed_data(input_data)
|
||||
|
||||
# Assert
|
||||
assert result['timestamp'] == {}
|
||||
assert result['variable'] == {}
|
||||
assert result['value'] == {}
|
||||
assert result['model_id'] == {}
|
||||
gates_activity.info.assert_called()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_format_default_prediction(gates_activity):
|
||||
# Arrange
|
||||
@@ -603,6 +677,40 @@ async def test_format_retrain_report(gates_activity):
|
||||
assert result['mlflow_experiment_id'] == {0: 'test_mlflow_experiment_id'}
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_format_retrain_report_failure(gates_activity):
|
||||
# Arrange
|
||||
input_data = {
|
||||
**metadata,
|
||||
'experiment_response': {
|
||||
'success': False,
|
||||
'timestamp': '2023-05-26 11:12:27',
|
||||
'message': 'failure',
|
||||
},
|
||||
'update_report': {
|
||||
'version': '1.0.0',
|
||||
'mlflow_run_id': 'test_mlflow_run_id',
|
||||
'mlflow_experiment_id': 'test_mlflow_experiment_id',
|
||||
},
|
||||
'model_id': 'test_model',
|
||||
'model_name': 'test_model',
|
||||
}
|
||||
|
||||
# Act
|
||||
result = await gates_activity.format_retrain_report(input_data)
|
||||
|
||||
# Assert
|
||||
assert result['model_id'] == {0: 'test_model'}
|
||||
assert result['model_name'] == {0: 'test_model'}
|
||||
assert result['timestamp'] == {0: '2023-05-26 11:12:27'}
|
||||
assert result['status'] == {0: 'failure'}
|
||||
assert 'version' not in result
|
||||
assert 'mlflow_run_id' not in result
|
||||
assert 'mlflow_experiment_id' not in result
|
||||
gates_activity.info.assert_called()
|
||||
gates_activity.debug.assert_called()
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_get_last_timestamp_with_data(gates_activity):
|
||||
# Arrange
|
||||
@@ -742,3 +850,41 @@ async def test_write_metrics(mock_metrics, gates_activity):
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.gates.metrics')
|
||||
async def test_write_metrics_with_none_opc_response_time(mock_metrics, gates_activity):
|
||||
"""Test write_metrics method with None response_time in opc_metrics."""
|
||||
input_data = {
|
||||
**metadata,
|
||||
'prediction': {
|
||||
'prediction': [1],
|
||||
'prediction_confidence': [0.9],
|
||||
'response_time': [0.1],
|
||||
},
|
||||
'opc_metrics': {'server1': {'tag1': 0.1, 'tag2': None}},
|
||||
}
|
||||
await gates_activity.write_metrics(input_data)
|
||||
|
||||
# Verify that metrics for tag1 are emitted
|
||||
gates_activity.emit_metric.assert_any_call(
|
||||
metric_object=mock_metrics.PREDICTION_OPC_WRITING_RESPONSE_TIME_MONITOR,
|
||||
method='observe',
|
||||
tags={
|
||||
'pod_id': gates_activity.pod_id,
|
||||
'model_name': metadata['metadata']['model_name'],
|
||||
'workflow_name': metadata['metadata']['workflow_name'],
|
||||
'opc_server_id': 'server1',
|
||||
'tag': 'tag1',
|
||||
},
|
||||
value=0.1,
|
||||
)
|
||||
|
||||
# Verify that metrics for tag2 (with None response_time) are NOT emitted
|
||||
calls = [
|
||||
c
|
||||
for c in gates_activity.emit_metric.call_args_list
|
||||
if len(c[1].get('tags', {})) > 0 and c[1]['tags'].get('tag') == 'tag2'
|
||||
]
|
||||
assert len(calls) == 0, 'Metrics should not be emitted for None response_time'
|
||||
|
||||
@@ -76,6 +76,11 @@ def mlflow(mock_minio_repository, mock_mlflow_repository):
|
||||
mlflow.send_notification = MagicMock()
|
||||
mlflow.emit_metric = AsyncMock()
|
||||
mlflow.send_notification_async = AsyncMock()
|
||||
mlflow.error = MagicMock()
|
||||
mlflow.debug = MagicMock()
|
||||
mlflow.info = MagicMock()
|
||||
mlflow.warning = MagicMock()
|
||||
mlflow.critical = MagicMock()
|
||||
|
||||
return mlflow
|
||||
|
||||
@@ -488,3 +493,87 @@ async def test_update_production_model_error(mlflow):
|
||||
)
|
||||
else:
|
||||
raise AssertionError('No exception raised')
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch('laborious.activities.mlflow.to_datetime')
|
||||
async def test_get_reference_data_success(mock_to_datetime, mlflow):
|
||||
# Arrange
|
||||
input_data = {
|
||||
**metadata,
|
||||
'model_name': 'test_model',
|
||||
}
|
||||
|
||||
# Mock reference data DataFrame
|
||||
mock_reference_data = MagicMock()
|
||||
mock_reference_data.__getitem__.return_value = MagicMock()
|
||||
mock_to_datetime.return_value.dt.strftime.return_value = MagicMock()
|
||||
mock_reference_data.to_dict.return_value = [
|
||||
{'timestamp': '2023-05-26 11:12:27', 'value': 1.0},
|
||||
{'timestamp': '2023-05-26 11:12:28', 'value': 2.0},
|
||||
]
|
||||
|
||||
mlflow.model_monitoring_repository.load_artifact_dataframe.return_value = mock_reference_data
|
||||
|
||||
# Act
|
||||
result = await mlflow.get_reference_data(input_data)
|
||||
|
||||
# Assert
|
||||
mlflow.model_monitoring_repository.load_artifact_dataframe.assert_called_once_with(
|
||||
model_name='test_model',
|
||||
artifact_path='evaluation_data.csv',
|
||||
metadata=metadata['metadata'],
|
||||
)
|
||||
mock_to_datetime.assert_called_once_with(mock_reference_data.__getitem__.return_value)
|
||||
|
||||
mock_reference_data.to_dict.assert_called_once_with(orient='records')
|
||||
assert result == mock_reference_data.to_dict.return_value
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_get_reference_data_not_found(mlflow):
|
||||
# Arrange
|
||||
input_data = {
|
||||
**metadata,
|
||||
'model_name': 'test_model',
|
||||
}
|
||||
|
||||
mlflow.model_monitoring_repository.load_artifact_dataframe.return_value = None
|
||||
|
||||
# Act
|
||||
result = await mlflow.get_reference_data(input_data)
|
||||
|
||||
# Assert
|
||||
mlflow.model_monitoring_repository.load_artifact_dataframe.assert_called_once_with(
|
||||
model_name='test_model',
|
||||
artifact_path='evaluation_data.csv',
|
||||
metadata=metadata['metadata'],
|
||||
)
|
||||
mlflow.warning.assert_called_once_with(
|
||||
'Reference data not found for model test_model', metadata['metadata']
|
||||
)
|
||||
assert result is None
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
async def test_get_reference_data_exception(mlflow):
|
||||
# Arrange
|
||||
input_data = {
|
||||
**metadata,
|
||||
'model_name': 'test_model',
|
||||
}
|
||||
|
||||
mlflow.model_monitoring_repository.load_artifact_dataframe.side_effect = Exception(
|
||||
'Error loading artifact'
|
||||
)
|
||||
|
||||
# Act & Assert
|
||||
with raises(Exception) as e:
|
||||
await mlflow.get_reference_data(input_data)
|
||||
|
||||
assert str(e.value) == 'Error loading artifact'
|
||||
mlflow.model_monitoring_repository.load_artifact_dataframe.assert_called_once_with(
|
||||
model_name='test_model',
|
||||
artifact_path='evaluation_data.csv',
|
||||
metadata=metadata['metadata'],
|
||||
)
|
||||
|
||||
@@ -125,6 +125,156 @@ async def test_run_none_path_flag(workflow_mock, format_and_export_prediction):
|
||||
assert workflow_mock.execute_local_activity_method.call_count == 1
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch(
|
||||
'laborious.workflows.sub_workflows.format_and_export_prediction.workflow',
|
||||
new_callable=AsyncMock,
|
||||
)
|
||||
async def test_run_none_path_flag_with_transformed_data(
|
||||
workflow_mock, format_and_export_prediction
|
||||
):
|
||||
# Arrange
|
||||
input_data = {
|
||||
'metadata': metadata,
|
||||
'path_flag': None,
|
||||
'data': {'test': 'data'},
|
||||
'transformed_data': {'transformed': 'data'},
|
||||
'timestamp': '2021-01-01',
|
||||
'model_id': 1,
|
||||
'prediction_confidence': 0.9,
|
||||
'schema': 'test_schema',
|
||||
'table_name': 'test_table',
|
||||
'transform_table_name': 'test_transform_table',
|
||||
'opc_servers': ['test_server'],
|
||||
'opc_output_config': {'test': 'config'},
|
||||
'prediction_store_policy': 'lts:1',
|
||||
}
|
||||
|
||||
prediction_data = MagicMock()
|
||||
opc_metrics = MagicMock()
|
||||
transformed_data = MagicMock()
|
||||
|
||||
workflow_mock.execute_local_activity_method.side_effect = [
|
||||
prediction_data, # format_prediction
|
||||
transformed_data, # format_transformed_data
|
||||
]
|
||||
|
||||
write_transformed_handler = AsyncMock()
|
||||
workflow_mock.start_activity_method.return_value = write_transformed_handler
|
||||
workflow_mock.execute_activity_method.side_effect = [
|
||||
(prediction_data, opc_metrics), # write_opc_data
|
||||
MagicMock(), # export_data_to_postgres (prediction)
|
||||
MagicMock(), # write_metrics
|
||||
]
|
||||
|
||||
# Act
|
||||
await format_and_export_prediction.run(input_data)
|
||||
|
||||
# Assert - format_prediction call
|
||||
workflow_mock.execute_local_activity_method.assert_has_calls(
|
||||
[
|
||||
call(
|
||||
Activities.format_prediction,
|
||||
{
|
||||
'data': input_data['data'],
|
||||
'timestamp': input_data['timestamp'],
|
||||
'model_id': input_data['model_id'],
|
||||
'prediction_confidence': input_data['prediction_confidence'],
|
||||
'prediction_store_policy': input_data['prediction_store_policy'],
|
||||
**metadata,
|
||||
},
|
||||
retry_policy=ANY,
|
||||
start_to_close_timeout=ANY,
|
||||
),
|
||||
call(
|
||||
Activities.format_transformed_data,
|
||||
{
|
||||
'data': input_data['transformed_data'],
|
||||
'model_id': input_data['model_id'],
|
||||
**metadata,
|
||||
},
|
||||
retry_policy=ANY,
|
||||
start_to_close_timeout=ANY,
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
# Assert - start_activity_method for transformed data export
|
||||
workflow_mock.start_activity_method.assert_called_once_with(
|
||||
Activities.export_data_to_postgres,
|
||||
{
|
||||
'schema': input_data['schema'],
|
||||
'table_name': input_data['transform_table_name'],
|
||||
'data': transformed_data,
|
||||
'timestamp_conversion': {
|
||||
'column': 'timestamp',
|
||||
'format': DATETIME_FORMAT_WITH_TZ,
|
||||
},
|
||||
**metadata,
|
||||
},
|
||||
retry_policy=ANY,
|
||||
start_to_close_timeout=ANY,
|
||||
)
|
||||
|
||||
# Assert - write_opc_data call
|
||||
workflow_mock.execute_activity_method.assert_has_calls(
|
||||
[
|
||||
call(
|
||||
Activities.write_opc_data,
|
||||
{
|
||||
'opc_output_config': input_data['opc_output_config'],
|
||||
'data': prediction_data,
|
||||
**metadata,
|
||||
},
|
||||
retry_policy=ANY,
|
||||
start_to_close_timeout=ANY,
|
||||
)
|
||||
]
|
||||
)
|
||||
|
||||
# Assert - export_data_to_postgres for prediction call
|
||||
workflow_mock.execute_activity_method.assert_has_calls(
|
||||
[
|
||||
call(
|
||||
Activities.export_data_to_postgres,
|
||||
{
|
||||
'schema': input_data['schema'],
|
||||
'table_name': input_data['table_name'],
|
||||
'data': prediction_data,
|
||||
'timestamp_conversion': {
|
||||
'column': 'timestamp',
|
||||
'format': DATETIME_FORMAT_WITH_TZ,
|
||||
},
|
||||
**metadata,
|
||||
},
|
||||
retry_policy=ANY,
|
||||
start_to_close_timeout=ANY,
|
||||
)
|
||||
]
|
||||
)
|
||||
|
||||
# Assert - write_metrics call
|
||||
workflow_mock.execute_activity_method.assert_has_calls(
|
||||
[
|
||||
call(
|
||||
Activities.write_metrics,
|
||||
{
|
||||
**metadata,
|
||||
'prediction': prediction_data,
|
||||
'opc_metrics': opc_metrics,
|
||||
},
|
||||
retry_policy=ANY,
|
||||
start_to_close_timeout=ANY,
|
||||
)
|
||||
]
|
||||
)
|
||||
|
||||
# Assert - verify counts
|
||||
assert workflow_mock.execute_activity_method.call_count == 3
|
||||
assert workflow_mock.execute_local_activity_method.call_count == 2
|
||||
assert workflow_mock.start_activity_method.call_count == 1
|
||||
|
||||
|
||||
@mark.asyncio
|
||||
@patch(
|
||||
'laborious.workflows.sub_workflows.format_and_export_prediction.workflow',
|
||||
|
||||
Reference in New Issue
Block a user