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