SIENTIAPDE-1241: refactor train_model workflow due to I/O errors.

This commit is contained in:
Bruno Domingues
2025-10-22 15:37:56 -03:00
parent f2a1c88ff3
commit 5789a13023
31 changed files with 37878 additions and 6097 deletions

View File

@@ -10,7 +10,6 @@ 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
@@ -18,8 +17,10 @@ with workflow.unsafe.imports_passed_through():
from sientia_do.observability.logger import Logger
from sientia_do.temporal.activities.base import BaseActivity
from model_manager.utils.exceptions import ModelTrainingError
from model_manager.utils.models.train_model_params import TrainModelParams
from model_manager.utils.models.train_model_result import TrainModelResult
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
@@ -39,6 +40,8 @@ class Training(BaseActivity):
def __init__(
self,
model_repository: ModelRepository,
storage_repository: StorageRepository,
logger: Logger,
notification_handler: NotificationHandler,
):
@@ -49,8 +52,10 @@ class Training(BaseActivity):
logger: Logger instance for observability
notification_handler: Handler for sending notifications
"""
BaseActivity.__init__(self, logger, notification_handler, set_error_counter=True)
super().__init__(logger, notification_handler, 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:
@@ -71,27 +76,12 @@ class Training(BaseActivity):
Raises:
ValueError, TypeError, KeyError: If validation fails (after sending notification)
Example:
result = await validate_train_params({
'metadata': {'workflow_id': 'train-123'},
'experiment_run_id': 456,
'target_variable': 'price',
'variable_columns': ['feature1', 'feature2'],
'train_size': 80,
# ... other required fields at same level
})
# Returns: TrainModelParams(...)
"""
metadata = input_data.get('metadata', {})
try:
self.info('Validating training parameters', metadata)
# Step 1: Convert input_data to TrainModelParams (validates types and required fields)
train_params = TrainModelParams.from_dict(input_data)
# Step 2: Validate business rules (ranges, consistency, etc.)
train_params.validate_business_rules()
self.info(
@@ -102,12 +92,10 @@ class Training(BaseActivity):
)
return train_params
except (ValueError, TypeError, KeyError) as e:
error_msg = f'Error validating training parameters: {str(e)}'
trace = traceback.format_exc()
# Send notification (MongoDB)
self.send_notification(
metadata=metadata,
notification_id='VALIDATE_TRAIN_PARAMS_ERROR',
@@ -117,14 +105,11 @@ class Training(BaseActivity):
attachment_content=trace,
)
# Log error with metadata
self.error(trace, metadata=metadata)
# Re-raise exception to stop workflow
raise
@activity.defn(name='train_model')
async def train_model(self, input_data: dict[str, Any]) -> TrainModelResult:
async def train_model(self, input_data: dict[str, Any]) -> dict[str, str | None]:
"""
Train a machine learning model.
@@ -156,51 +141,42 @@ class Training(BaseActivity):
# Returns: TrainModelResult(...)
"""
metadata = input_data.get('metadata', {})
uploaded_file = input_data['uploaded_file']
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
try:
# Validate uploaded_file is BytesIO
if not isinstance(uploaded_file, BytesIO):
raise ValueError(f'uploaded_file must be BytesIO, got {type(uploaded_file)}')
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)
# Validate train_params is TrainModelParams
if not isinstance(train_params, TrainModelParams):
raise ValueError(f'train_params must be TrainModelParams, got {type(train_params)}')
train_result = self.training_repository.after_train_calculation(
train_params, train_result
)
self.info(
f'Starting model training for target: {train_params.target_variable}',
metadata,
)
# 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 final_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
target = (
train_params.target_variable
if hasattr(train_params, 'target_variable')
else 'unknown'
error_msg = (
'Error training model - '
f'model_trained={model_trained}, model_saved={model_saved}, '
f'error: {str(e)}'
)
error_msg = f'Error training model - Target: {target}, Error: {str(e)}'
trace = traceback.format_exc()
# Send notification (MongoDB)
self.send_notification(
metadata=metadata,
notification_id='TRAIN_MODEL_ERROR',
@@ -210,8 +186,39 @@ class Training(BaseActivity):
attachment_content=trace,
)
# Log error with metadata
self.error(trace, metadata=metadata)
# Re-raise exception to stop workflow
raise ModelTrainingError(
model_trained=model_trained,
model_saved=model_saved,
) from e
@activity.defn(name='cleanup_resources')
async def cleanup_resources(self, input_data: dict[str, Any]) -> None:
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', '')
try:
self.model_repository.cleanup_run_directory(run_dir)
self.storage_repository.delete_file(bucket_name, file_name)
except Exception as e: # noqa: BLE001
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,
)
self.error(trace, metadata=metadata)
raise