feat: update training workflow and repository management

- Replaced synchronous MinIO repository calls with asynchronous counterparts in the Training class for improved performance.
- Enhanced logging throughout the training process to provide better insights into model metadata loading, parameter validation, and training execution.
- Updated the train_test_split function to enforce DataFrame input type, ensuring consistency in data handling.
- Removed the deprecated model_repository.py file to streamline the codebase.
- Adjusted cleanup schedule logic to improve error handling and logging during schedule reconciliation.
- Updated tests to reflect changes in the training workflow and repository interactions.
This commit is contained in:
vitor-aignosi
2026-04-09 12:09:52 -03:00
parent 0ae03b246f
commit 526edcb50e
14 changed files with 114 additions and 503 deletions

View File

@@ -18,7 +18,7 @@ with workflow.unsafe.imports_passed_through():
from sientia_do.observability.logger import Logger
from sientia_do.observability.metrics_controller import MetricsController
from sientia_do.observability.sientia_monitoring import SientiaMonitoring
from sientia_do.repository.minio_repository import MinioRepository
from sientia_do.repository.minio_repository_sync import MinioRepository
from sientia_model.model_repository.mlflow_repository import SientiaMLflowRepository
from sientia_model.model_repository.plugin_store import PluginStore
@@ -78,14 +78,18 @@ class Training(SientiaMonitoring):
"""
metadata = input_data.get('metadata', {})
self.info(f'Loading model metadata for {input_data}', metadata)
try:
train_params = TrainModelParams.from_dict(input_data)
model_metadata = self.plugin_store.get_model_index(
model_name=train_params.model_name,
model_type=train_params.model_type,
metadata=metadata,
)
train_params.model_metadata = model_metadata
self.info(f'Model metadata loaded successfully for {input_data}', metadata)
self.debug(f'Model metadata: {model_metadata}', metadata)
return train_params.to_dict()
except Exception as exc:
trace = traceback.format_exc()
@@ -120,6 +124,9 @@ class Training(SientiaMonitoring):
Exception: If validation fails (after sending notification)
"""
metadata = input_data.get('metadata', {})
self.info(f'Validating training parameters for {input_data}', metadata)
try:
train_params = TrainModelParams.from_dict(input_data)
@@ -132,6 +139,10 @@ class Training(SientiaMonitoring):
metadata,
)
self.debug(
f'Training parameters validated successfully: {train_params.to_dict()}', metadata
)
return train_params.to_dict()
except Exception as e:
error_msg = f'Error validating training parameters: {str(e)}'
@@ -173,9 +184,15 @@ class Training(SientiaMonitoring):
metadata = input_data.get('metadata')
train_params = TrainModelParams.from_dict(input_data['train_params'])
self.info('Starting train_model process', metadata)
try:
# Download training file bytes from MinIO
train_bytes = self.minio_repository.download_file_sync(
self.info(
f'Downloading training file from MinIO for {train_params.file_name}', metadata
)
train_bytes = self.minio_repository.download_file(
object_name=train_params.file_name,
bucket=train_params.bucket_name,
metadata=metadata,
@@ -185,12 +202,14 @@ class Training(SientiaMonitoring):
val_bytes: bytes | None = None
validation_name = train_params.val_file_name
if validation_name is not None:
val_bytes = self.minio_repository.download_file_sync(
self.info(f'Downloading validation file from MinIO for {validation_name}', metadata)
val_bytes = self.minio_repository.download_file(
object_name=validation_name,
bucket=train_params.bucket_name,
metadata=metadata,
)
self.info(f'Preparing training data for {train_params.file_name}', metadata)
train_result = self.data_manager_repository.prepare_training_data(
train_file_bytes=train_bytes,
validation_file_bytes=val_bytes,
@@ -198,8 +217,9 @@ class Training(SientiaMonitoring):
metadata=metadata,
)
self.info(f'Getting model wrapper for {train_params.model_type}', metadata)
wrapper = self.plugin_store.get_model(
model_name=train_params.model_name,
model_type=train_params.model_type,
force_download=False,
opt_params=train_params.opt_params or {},
model_kwargs=train_params.model_kwargs or {},
@@ -207,6 +227,7 @@ class Training(SientiaMonitoring):
metadata=metadata,
)
self.info(f'Training model for {train_params.model_type}', metadata)
train_data = train_result.train_data
val_data = train_result.val_data
@@ -216,6 +237,10 @@ class Training(SientiaMonitoring):
target=train_params.target_variable,
)
self.info(
f'Generating predictions using the trained wrapper for {train_params.model_type}',
metadata,
)
# Generate predictions using the trained wrapper
transformed_train, _ = wrapper.transform(train_data)
transformed_val, _ = wrapper.transform(val_data)
@@ -229,11 +254,13 @@ class Training(SientiaMonitoring):
train_result.y_train_pred = y_train_pred_df
train_result.y_pred = y_val_pred_df
self.info(f'Computing regression metrics for {train_params.model_type}', metadata)
train_result = self.data_manager_repository.compute_regression_metrics(
train_result,
wrapper,
)
self.info(f'Starting MLflow run for {train_params.model_type}', metadata)
with self.mlflow_repository.start_run(
model_name=train_params.model_name,
run_name=None,
@@ -273,6 +300,7 @@ class Training(SientiaMonitoring):
wrapper: Any,
metadata: dict[str, Any] | None,
) -> None:
self.info(f'Generating report for {train_params.model_type}', metadata)
train_result = self.data_manager_repository.generate_report(
train_result,
metadata=metadata,
@@ -285,7 +313,9 @@ class Training(SientiaMonitoring):
):
raise ValueError('Report path, train data path, or test data path is not set')
self.info(f'Storing model for {train_params.model_type}', metadata)
wrapper.store_model(name=train_params.model_name)
self.info(f'Logging artifacts for {train_params.model_type}', metadata)
mlflow.log_artifact(train_result.report_path)
mlflow.log_artifact(train_result.train_data_path)
mlflow.log_artifact(train_result.test_data_path)