SIENTIAPDE-1251: Implement ML model training activity and repository
This commit introduces the 'Training' activity and 'TrainingRepository' for handling ML model training operations within the Model Manager system. - Added model_manager/activities/training.py for the Training activity, which extends BaseActivity and integrates with Temporal workflows. - Added model_manager/utils/repository/training_repository.py for the TrainingRepository, which encapsulates the core training logic. - Updated model_manager/activities/activities.py to include the Training activity in the main activities orchestrator. - Updated README.md to document the new 'Training' component. - Added unit tests for the new activity and repository.
This commit is contained in:
@@ -10,9 +10,10 @@ with workflow.unsafe.imports_passed_through():
|
||||
from model_manager.activities.gates import Gates
|
||||
from model_manager.activities.minio import MinIO
|
||||
from model_manager.activities.mlflow import MLFlow
|
||||
from model_manager.activities.training import Training
|
||||
|
||||
|
||||
class Activities(ExperimentTracking, MLFlow, MinIO, Gates):
|
||||
class Activities(ExperimentTracking, MLFlow, MinIO, Gates, Training):
|
||||
"""
|
||||
Main activities orchestrator for the Model Manager system.
|
||||
|
||||
@@ -25,6 +26,7 @@ class Activities(ExperimentTracking, MLFlow, MinIO, Gates):
|
||||
- MLFlow: Model inference and transformation operations
|
||||
- MinIO: Object storage operations (file upload/download/delete)
|
||||
- Gates: Data quality validation and filtering mechanisms
|
||||
- Training: ML model training operations (extends BaseActivity)
|
||||
|
||||
Attributes:
|
||||
postgres_config (dict): PostgreSQL connection configuration
|
||||
@@ -102,6 +104,8 @@ class Activities(ExperimentTracking, MLFlow, MinIO, Gates):
|
||||
|
||||
Gates.__init__(self, logger=logger, notification_handler=notification_handler)
|
||||
|
||||
Training.__init__(self, logger=logger, notification_handler=notification_handler)
|
||||
|
||||
async def shutdown(self):
|
||||
"""
|
||||
Gracefully shutdown all activities and clean up resources.
|
||||
|
||||
166
model_manager/activities/training.py
Normal file
166
model_manager/activities/training.py
Normal file
@@ -0,0 +1,166 @@
|
||||
"""
|
||||
Training activities for ML model training operations.
|
||||
|
||||
This module provides activities for training machine learning models.
|
||||
The activity extends BaseActivity and receives pre-downloaded files
|
||||
to return success/failure status without raising exceptions.
|
||||
"""
|
||||
|
||||
from temporalio import activity, workflow
|
||||
|
||||
with workflow.unsafe.imports_passed_through():
|
||||
import traceback
|
||||
from io import BytesIO
|
||||
from typing import Any
|
||||
|
||||
from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler
|
||||
from sientia_do.notifications.models import NotificationLevel
|
||||
from sientia_do.observability.logger import Logger
|
||||
from sientia_do.temporal.activities.base import BaseActivity
|
||||
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
from model_manager.utils.repository.training_repository import TrainingRepository
|
||||
|
||||
|
||||
class Training(BaseActivity):
|
||||
"""
|
||||
Activity for ML model training operations.
|
||||
|
||||
This activity extends BaseActivity and handles machine learning model
|
||||
training with comprehensive error handling. It receives pre-downloaded
|
||||
files from the workflow and returns success/failure status without
|
||||
raising exceptions.
|
||||
|
||||
Attributes:
|
||||
logger (Logger): Logger instance for observability (inherited from BaseActivity)
|
||||
notification_handler (NotificationHandler): Handler for sending notifications (inherited)
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
logger: Logger,
|
||||
notification_handler: NotificationHandler,
|
||||
):
|
||||
"""
|
||||
Initialize Training activity.
|
||||
|
||||
Args:
|
||||
logger: Logger instance for observability
|
||||
notification_handler: Handler for sending notifications
|
||||
"""
|
||||
BaseActivity.__init__(self, logger, notification_handler, set_error_counter=True)
|
||||
self.training_repository = TrainingRepository(logger)
|
||||
|
||||
@activity.defn(name='train_model')
|
||||
async def train_model(self, input_data: dict[str, Any]) -> dict[str, Any]:
|
||||
"""
|
||||
Train a machine learning model with comprehensive error handling.
|
||||
|
||||
This activity orchestrates the complete ML training pipeline:
|
||||
1. Validates input parameters
|
||||
2. Trains the model using TrainingRepository
|
||||
3. Performs post-training calculations
|
||||
4. Returns success/failure status with results or error message
|
||||
|
||||
The activity does NOT raise exceptions on failure - it catches all errors,
|
||||
sends notifications, and returns a failure status. This allows the workflow
|
||||
to handle the error gracefully and update the database accordingly.
|
||||
|
||||
Args:
|
||||
input_data: Configuration for model training operation
|
||||
Required keys:
|
||||
- metadata (dict): Workflow execution metadata
|
||||
- uploaded_file (BytesIO): Training data file (already downloaded from MinIO)
|
||||
- train_params (dict): Training parameters (converted to TrainModelParams)
|
||||
|
||||
Returns:
|
||||
dict: Training result with the following structure:
|
||||
{
|
||||
'success': bool, # True if training succeeded, False otherwise
|
||||
'result': TrainModelResult | None, # Training result if success=True
|
||||
'error_message': str | None # Error message if success=False
|
||||
}
|
||||
|
||||
Example:
|
||||
# Successful training
|
||||
result = await train_model({
|
||||
'metadata': {'workflow_id': 'train-123', 'experiment_run_id': 456},
|
||||
'uploaded_file': BytesIO(csv_data),
|
||||
'train_params': {
|
||||
'experiment_run_id': 456,
|
||||
'target_variable': 'price',
|
||||
'variable_columns': ['feature1', 'feature2'],
|
||||
'train_size': 80,
|
||||
'shuffle': True,
|
||||
'use_scaler': True,
|
||||
# ... other TrainModelParams fields
|
||||
}
|
||||
})
|
||||
# Returns: {'success': True, 'result': TrainModelResult(...), 'error_message': None}
|
||||
|
||||
# Failed training
|
||||
# Returns: {'success': False, 'result': None, 'error_message': 'Error details...'}
|
||||
"""
|
||||
metadata = input_data.get('metadata', {})
|
||||
uploaded_file = input_data['uploaded_file']
|
||||
train_params_dict = input_data['train_params']
|
||||
|
||||
try:
|
||||
self.info(
|
||||
f'Starting model training for target: {train_params_dict.get("target_variable")}',
|
||||
metadata,
|
||||
)
|
||||
|
||||
# Convert dict to TrainModelParams
|
||||
train_params = TrainModelParams.from_dict(train_params_dict)
|
||||
|
||||
# Validate uploaded_file is BytesIO
|
||||
if not isinstance(uploaded_file, BytesIO):
|
||||
raise ValueError(f'uploaded_file must be BytesIO, got {type(uploaded_file)}')
|
||||
|
||||
# Step 1: Train the model
|
||||
self.info('Training model with TrainingRepository', metadata)
|
||||
train_result = self.training_repository.train(uploaded_file, train_params)
|
||||
|
||||
# Step 2: Perform post-training calculations
|
||||
self.info('Performing post-training calculations', metadata)
|
||||
final_result = self.training_repository.after_train_calculation(
|
||||
train_params, train_result
|
||||
)
|
||||
|
||||
self.info(
|
||||
f'Model training completed successfully - '
|
||||
f'MSE: {final_result.mse_val}, MAE: {final_result.mae_val}, R²: {final_result.r2_val}',
|
||||
metadata,
|
||||
)
|
||||
|
||||
return {
|
||||
'success': True,
|
||||
'result': final_result,
|
||||
'error_message': None,
|
||||
}
|
||||
|
||||
except Exception as e: # noqa: BLE001
|
||||
error_msg = f'Error training model - Target: {train_params_dict.get("target_variable", "unknown")}, Error: {str(e)}'
|
||||
trace = traceback.format_exc()
|
||||
|
||||
# Send notification (MongoDB)
|
||||
self.send_notification(
|
||||
metadata=metadata,
|
||||
notification_id='TRAIN_MODEL_ERROR',
|
||||
message=error_msg,
|
||||
block='train_model',
|
||||
level=NotificationLevel.ERROR,
|
||||
attachment_content=trace,
|
||||
)
|
||||
|
||||
# Log error with metadata
|
||||
self.error(trace, metadata=metadata)
|
||||
|
||||
# Return failure result (do NOT raise exception)
|
||||
# This allows workflow to update database with error status
|
||||
return {
|
||||
'success': False,
|
||||
'result': None,
|
||||
'error_message': str(e),
|
||||
}
|
||||
Reference in New Issue
Block a user