feat: enhance configuration and error handling in project setup

- Added new ignore rule for Ruff to allow temporary paths in tests.
- Introduced MyPy overrides for specific modules to ignore errors.
- Refactored `Cleanup` and `ExperimentTracking` classes to remove async keywords from methods, improving consistency in method signatures.
- Updated `Training` class methods to handle synchronous operations, enhancing performance and clarity.
- Adjusted `requirements.txt` to remove unnecessary Git dependency, streamlining project setup.
This commit is contained in:
vitor-aignosi
2026-04-07 10:25:17 -03:00
parent 6b1df7c3a7
commit 09ee92f100
21 changed files with 500 additions and 309 deletions

View File

@@ -12,6 +12,7 @@ with workflow.unsafe.imports_passed_through():
import traceback
from typing import Any
import mlflow
from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler
from sientia_do.notifications.models import NotificationLevel
from sientia_do.observability.logger import Logger
@@ -22,10 +23,8 @@ with workflow.unsafe.imports_passed_through():
from sientia_model.model_repository.plugin_store import PluginStore
from model_manager.utils.models.train_model_params import TrainModelParams
from model_manager.utils.repository.data_manager_repository import DataManagerRepository
from model_manager.utils.models.train_model_result import TrainModelResult
import mlflow
from model_manager.utils.repository.data_manager_repository import DataManagerRepository
class Training(SientiaMonitoring):
@@ -61,7 +60,7 @@ class Training(SientiaMonitoring):
self.minio_repository = minio_repository
@activity.defn(name='load_model_metadata')
async def load_model_metadata(self, input_data: dict[str, Any]) -> dict[str, Any]:
def load_model_metadata(self, input_data: dict[str, Any]) -> dict[str, Any]:
"""
Load model metadata/schemas from the model store.
@@ -90,7 +89,7 @@ class Training(SientiaMonitoring):
return train_params.to_dict()
except Exception as exc:
trace = traceback.format_exc()
await self.send_notification_async(
self.send_notification(
metadata=metadata,
notification_id='LOAD_MODEL_METADATA_ERROR',
message=f'Error loading model metadata: {str(exc)}',
@@ -99,9 +98,9 @@ class Training(SientiaMonitoring):
attachment_content=trace,
)
raise
@activity.defn(name='validate_train_params')
async def validate_train_params(self, input_data: dict[str, Any]) -> TrainModelParams:
def validate_train_params(self, input_data: dict[str, Any]) -> dict[str, Any]:
"""
Validate and convert training parameters from dict to TrainModelParams.
@@ -115,7 +114,7 @@ class Training(SientiaMonitoring):
- All TrainModelParams fields (experiment_run_id, target_variable, etc.)
Returns:
TrainModelParams: Validated and converted training parameters
dict[str, Any]: Validated and converted training parameters as dictionary
Raises:
Exception: If validation fails (after sending notification)
@@ -123,7 +122,7 @@ class Training(SientiaMonitoring):
metadata = input_data.get('metadata', {})
try:
train_params = TrainModelParams.from_dict(input_data)
train_params.validate_business_rules()
self.info(
@@ -133,12 +132,12 @@ class Training(SientiaMonitoring):
metadata,
)
return train_params
return train_params.to_dict()
except Exception as e:
error_msg = f'Error validating training parameters: {str(e)}'
trace = traceback.format_exc()
await self.send_notification_async(
self.send_notification(
metadata=metadata,
notification_id='VALIDATE_TRAIN_PARAMS_ERROR',
message=error_msg,
@@ -149,7 +148,7 @@ class Training(SientiaMonitoring):
raise
@activity.defn(name='train_model')
async def train_model(self, input_data: dict[str, Any]) -> dict[str, Any]:
def train_model(self, input_data: dict[str, Any]) -> dict[str, Any]:
"""
Train a machine learning model.
@@ -162,7 +161,7 @@ class Training(SientiaMonitoring):
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.
- train_params (dict): Training parameters.
Returns:
dict[str, Any]: Serializable summary (run identifiers, run_dir for cleanup, regression metrics).
@@ -172,14 +171,11 @@ class Training(SientiaMonitoring):
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)
train_params = TrainModelParams.from_dict(input_data['train_params'])
try:
# Download training file bytes from MinIO
train_bytes = await self.minio_repository.download_file(
train_bytes = self.minio_repository.download_file_sync(
object_name=train_params.file_name,
bucket=train_params.bucket_name,
metadata=metadata,
@@ -189,7 +185,7 @@ class Training(SientiaMonitoring):
val_bytes: bytes | None = None
validation_name = train_params.val_file_name
if validation_name is not None:
val_bytes = await self.minio_repository.download_file(
val_bytes = self.minio_repository.download_file_sync(
object_name=validation_name,
bucket=train_params.bucket_name,
metadata=metadata,
@@ -202,13 +198,13 @@ class Training(SientiaMonitoring):
metadata=metadata,
)
wrapper = await self.plugin_store.get_model(
wrapper = self.plugin_store.get_model(
model_name=train_params.model_name,
force_download=False,
opt_params=train_params.opt_params or {},
model_kwargs=train_params.model_kwargs or {},
data_model_kwargs=train_params.data_model_kwargs or {},
metadata=metadata
metadata=metadata,
)
train_data = train_result.train_data
@@ -238,7 +234,7 @@ class Training(SientiaMonitoring):
wrapper,
)
async with self.mlflow_repository.start_run(
with self.mlflow_repository.start_run(
model_name=train_params.model_name,
run_name=None,
experiment_name=f'{train_params.model_name}_experiment',
@@ -247,33 +243,19 @@ class Training(SientiaMonitoring):
) as run_info:
train_result.run_name = run_info.run_name
train_result.run_id = run_info.run_id
train_result = self.data_manager_repository.generate_report(
train_result,
metadata=metadata,
)
if train_result.report_path is None or train_result.train_data_path is None or train_result.test_data_path is None:
raise ValueError('Report path, train data path, or test data path is not set')
wrapper.store_model(name=train_params.model_name)
mlflow.log_artifact(train_result.report_path)
mlflow.log_artifact(train_result.train_data_path)
mlflow.log_artifact(train_result.test_data_path)
self._persist_training_artifacts(train_result, train_params, wrapper, metadata)
return {
'run_name': train_result.run_name,
'run_id': train_result.run_id,
'run_dir': train_result.run_dir
'run_dir': train_result.run_dir,
}
except Exception as e: # noqa: BLE001
error_msg = f'Error training model - error: {str(e)}'
trace = traceback.format_exc()
await self.send_notification_async(
self.send_notification(
metadata=metadata or {},
notification_id='TRAIN_MODEL_ERROR',
message=error_msg,
@@ -284,8 +266,32 @@ class Training(SientiaMonitoring):
raise e
def _persist_training_artifacts(
self,
train_result: TrainModelResult,
train_params: TrainModelParams,
wrapper: Any,
metadata: dict[str, Any] | None,
) -> None:
train_result = self.data_manager_repository.generate_report(
train_result,
metadata=metadata,
)
if (
train_result.report_path is None
or train_result.train_data_path is None
or train_result.test_data_path is None
):
raise ValueError('Report path, train data path, or test data path is not set')
wrapper.store_model(name=train_params.model_name)
mlflow.log_artifact(train_result.report_path)
mlflow.log_artifact(train_result.train_data_path)
mlflow.log_artifact(train_result.test_data_path)
@activity.defn(name='cleanup_resources')
async def cleanup_resources(self, input_data: dict[str, Any]) -> None:
def cleanup_resources(self, input_data: dict[str, Any]) -> None:
"""
Cleanup temporary resources created during training.
@@ -317,4 +323,3 @@ class Training(SientiaMonitoring):
)
raise