Files
sientia-dataops-model-manager/model_manager/activities/training.py

285 lines
10 KiB
Python

"""
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 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.observability.metrics_controller import MetricsController
from sientia_do.observability.sientia_monitoring import SientiaMonitoring
from model_manager.metrics import ACTIVITY_EXECUTION_TOTAL, WORKFLOW_EXECUTION_TOTAL
from model_manager.utils.exceptions import ModelTrainingError
from model_manager.utils.models.train_model_params import TrainModelParams
from model_manager.utils.repository.model_repository import ModelRepository
from model_manager.utils.repository.storage_repository import StorageRepository
from model_manager.utils.repository.training_repository import TrainingRepository
class Training(SientiaMonitoring):
"""
Activity for ML model training operations.
This activity extends SientiaMonitoring 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.
"""
def __init__(
self,
model_repository: ModelRepository,
storage_repository: StorageRepository,
logger: Logger,
notification_handler: NotificationHandler,
metrics_controller: MetricsController,
):
"""
Initialize Training activity.
Args:
logger: Logger instance for observability
notification_handler: Handler for sending notifications
"""
super().__init__(logger, notification_handler, metrics_controller, set_error_counter=True)
self.training_repository = TrainingRepository(logger)
self.model_repository = model_repository
self.storage_repository = storage_repository
@activity.defn(name='validate_train_params')
async def validate_train_params(self, input_data: dict[str, Any]) -> TrainModelParams:
"""
Validate and convert training parameters from dict to TrainModelParams.
This activity validates the input training parameters and converts them
to a TrainModelParams object.
Args:
input_data: Training parameters and metadata at the same level
Required keys:
- metadata (dict): Workflow execution metadata
- All TrainModelParams fields (experiment_run_id, target_variable, etc.)
Returns:
TrainModelParams: Validated and converted training parameters
Raises:
ValueError, TypeError, KeyError: If validation fails (after sending notification)
"""
metadata = input_data.get('metadata', {})
metrics_status = 'success'
try:
train_params = TrainModelParams.from_dict(input_data)
train_params.validate_business_rules()
self.info(
f'Training parameters validated successfully - '
f'Target: {train_params.target_variable}, '
f'Experiment: {train_params.experiment_name}',
metadata,
)
return train_params
except (ValueError, TypeError, KeyError) as e:
metrics_status = 'error'
error_msg = f'Error validating training parameters: {str(e)}'
trace = traceback.format_exc()
self.send_notification(
metadata=metadata,
notification_id='VALIDATE_TRAIN_PARAMS_ERROR',
message=error_msg,
block='validate_train_params',
level=NotificationLevel.ERROR,
attachment_content=trace,
)
raise
finally:
await self._emit_metrics(
metadata=metadata,
metrics_status=metrics_status,
activity_name='validate_train_params',
emit_workflow_metric=(metrics_status == 'error'),
)
@activity.defn(name='train_model')
async def train_model(self, input_data: dict[str, Any]) -> dict[str, str | None]:
"""
Train a machine learning model.
This activity orchestrates the ML training pipeline:
1. Validate input parameters.
2. Train the model via TrainingRepository.
3. Perform post-training calculations.
Args:
input_data: Training configuration containing:
- metadata (dict): Workflow execution metadata.
- uploaded_file (BytesIO): Training data already downloaded from MinIO.
- train_params (TrainModelParams | dict): Training parameters.
Returns:
dict: Keys `run_name` and `run_dir` when training and saving succeed.
Raises:
ValueError: If input validation fails.
Exception: If training fails (after sending notification).
"""
metadata = input_data.get('metadata', {})
train_params = input_data['train_params']
if isinstance(train_params, dict):
train_params = TrainModelParams.from_dict(train_params)
# type: ignore[assignment]
model_trained = False
model_saved = False
metrics_status = 'success'
try:
with self.storage_repository.fetch_file(
train_params.bucket_name, train_params.file_name
) as uploaded_file:
train_result = self.training_repository.train(uploaded_file, train_params)
train_result = self.training_repository.after_train_calculation(
train_params, train_result
)
model_trained = True
train_result = self.model_repository.save_model(train_result)
model_saved = True
return {
'run_name': train_result.run_name,
'run_dir': train_result.run_dir,
}
except Exception as e: # noqa: BLE001
metrics_status = 'error'
error_msg = (
'Error training model - '
f'model_trained={model_trained}, model_saved={model_saved}, '
f'error: {str(e)}'
)
trace = traceback.format_exc()
self.send_notification(
metadata=metadata,
notification_id='TRAIN_MODEL_ERROR',
message=error_msg,
block='train_model',
level=NotificationLevel.ERROR,
attachment_content=trace,
)
raise ModelTrainingError(
model_trained=model_trained,
model_saved=model_saved,
) from e
finally:
await self._emit_metrics(
metadata=metadata,
metrics_status=metrics_status,
activity_name='train_model',
emit_workflow_metric=(metrics_status == 'error'),
)
@activity.defn(name='cleanup_resources')
async def cleanup_resources(self, input_data: dict[str, Any]) -> None:
"""
Cleanup temporary resources created during training.
Args:
input_data: Cleanup configuration containing:
- metadata (dict): Workflow execution metadata.
- run_dir (str): Temporary directory to remove.
- bucket_name (str): MinIO bucket of the uploaded file.
- file_name (str): MinIO object key to delete.
Raises:
Exception: If cleanup fails (after sending notification).
"""
metadata = input_data.get('metadata', {})
run_dir = input_data.get('run_dir', '')
bucket_name = input_data.get('bucket_name', '')
file_name = input_data.get('file_name', '')
metrics_status = 'success'
try:
self.model_repository.cleanup_run_directory(run_dir)
self.storage_repository.delete_file(bucket_name, file_name)
except Exception as e: # noqa: BLE001
metrics_status = 'error'
error_msg = (
f'Error cleaning up resources - Run directory: {run_dir}, '
f'File: {bucket_name}/{file_name}, Error: {str(e)}'
)
trace = traceback.format_exc()
self.send_notification(
metadata=metadata,
notification_id='CLEANUP_RESOURCES_ERROR',
message=error_msg,
block='cleanup_resources',
level=NotificationLevel.ERROR,
attachment_content=trace,
)
raise
finally:
await self._emit_metrics(
metadata=metadata,
metrics_status=metrics_status,
activity_name='cleanup_resources',
emit_workflow_metric=True,
)
async def _emit_metrics(
self,
metadata: dict[str, Any],
metrics_status: str,
activity_name: str,
emit_workflow_metric: bool,
) -> None:
"""
Emit workflow and activity execution metrics.
Args:
metadata: Activity metadata containing pod_id and workflow_name
metrics_status: Execution status ('success' or 'error')
activity_name: Name of the activity being executed
"""
if emit_workflow_metric:
await self.emit_metric(
metric_object=WORKFLOW_EXECUTION_TOTAL,
tags={
'pod_id': metadata.get('pod_id'),
'workflow_name': metadata.get('workflow_name'),
'status': metrics_status,
},
)
await self.emit_metric(
metric_object=ACTIVITY_EXECUTION_TOTAL,
tags={
'pod_id': metadata.get('pod_id'),
'activity_name': activity_name,
'status': metrics_status,
},
)