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:
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user