feat: enhance training workflow with model metadata loading and refactor data handling

- Introduced a new activity to load model metadata from the model store.
- Refactored training logic to utilize new model metadata and improved parameter handling.
- Updated the `TrainModelParams` class to include additional fields for model configuration.
- Replaced deprecated utility functions with a custom train-test split implementation.
- Removed unused utility functions and cleaned up the data manager repository.
- Adjusted experiment tracking to include model-specific metadata in notifications.
This commit is contained in:
vitor-aignosi
2026-03-24 14:39:28 -03:00
parent cf5111e520
commit 342a02d6f7
13 changed files with 242 additions and 465 deletions

View File

@@ -246,7 +246,7 @@ class ExperimentTracking(Postgres):
ValueError: If required parameters are missing for the update type
RuntimeError: If update operation fails
"""
metadata = input_data.get('metadata', {})
metadata = input_data.get('metadata')
experiment_run_id = input_data['experiment_run_id']
update_type = input_data['update_type']
status = input_data.get('status')
@@ -272,8 +272,8 @@ class ExperimentTracking(Postgres):
error_msg = f'Error updating experiment run - ID: {experiment_run_id}, Status: {status}, Error: {str(e)}'
trace = traceback.format_exc()
self.send_notification(
metadata=metadata,
await self.send_notification_async(
metadata=metadata or {},
notification_id='UPDATE_EXPERIMENT_RUN_ERROR',
message=error_msg,
block='update_experiment_run',

View File

@@ -21,6 +21,7 @@ with workflow.unsafe.imports_passed_through():
from sientia_do.repository.minio_repository import MinioRepository
from sientia_model.model_repository.mlflow_repository import SientiaMLflowRepository
from sientia_model.model_repository.plugin_store import PluginStore
from sientia_model.wrappers.sientia_model import SientiaModel
from model_manager.metrics import ACTIVITY_EXECUTION_TOTAL, WORKFLOW_EXECUTION_TOTAL
from model_manager.utils.exceptions import ModelTrainingError
@@ -60,6 +61,46 @@ class Training(SientiaMonitoring):
self.plugin_store = plugin_store
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]:
"""
Load model metadata/schemas from the model store.
This activity is responsible for fetching model metadata/schemas from the
model store index and extracting a serializable `model_metadata` dict that
`TrainModelParams.validate_business_rules()` depends on.
Args:
input_data: Workflow input at the same level as `validate_train_params`,
including at least `model_name` and the fields required by
`TrainModelParams.from_dict` to build wrapper kwargs.
Return:
dict[str, Any]: Updated `input_data` containing `input_data['model_metadata']`.
"""
metadata = input_data.get('metadata', {})
try:
train_params = TrainModelParams.from_dict(input_data)
model_metadata = self.plugin_store.get_model_index(
model_name=train_params.model_name,
metadata=metadata,
)
train_params.model_metadata = model_metadata
return train_params.to_dict()
except Exception as exc:
trace = traceback.format_exc()
await self.send_notification_async(
metadata=metadata,
notification_id='LOAD_MODEL_METADATA_ERROR',
message=f'Error loading model metadata: {str(exc)}',
block='load_model_metadata',
level=NotificationLevel.ERROR,
attachment_content=trace,
)
raise
@activity.defn(name='validate_train_params')
async def validate_train_params(self, input_data: dict[str, Any]) -> TrainModelParams:
"""
@@ -81,10 +122,9 @@ class Training(SientiaMonitoring):
ValueError, TypeError, KeyError: If validation fails (after sending notification)
"""
metadata = input_data.get('metadata', {})
metrics_status = 'success'
try:
train_params = TrainModelParams.from_dict(input_data)
train_params.validate_business_rules()
self.info(
@@ -96,11 +136,10 @@ class Training(SientiaMonitoring):
return train_params
except (ValueError, TypeError, KeyError) as e:
metrics_status = 'error'
error_msg = f'Error validating training parameters: {str(e)}'
trace = traceback.format_exc()
self.send_notification(
await self.send_notification_async(
metadata=metadata,
notification_id='VALIDATE_TRAIN_PARAMS_ERROR',
message=error_msg,
@@ -109,13 +148,6 @@ class Training(SientiaMonitoring):
attachment_content=trace,
)
raise
finally:
await self._emit_metrics(
metadata=metadata,
metrics_status=metrics_status,
activity_name='validate_train_params',
emit_workflow_metric=(metrics_status == 'error'),
)
@activity.defn(name='train_model')
async def train_model(self, input_data: dict[str, Any]) -> dict[str, str | None]:
@@ -146,10 +178,8 @@ class Training(SientiaMonitoring):
if isinstance(train_params, dict):
train_params = TrainModelParams.from_dict(train_params)
model_trained = False
model_saved = False
metrics_status = 'success'
try:
# Download training file bytes from MinIO
@@ -161,7 +191,7 @@ class Training(SientiaMonitoring):
# Download optional validation file bytes from the same bucket
val_bytes: bytes | None = None
validation_name = getattr(train_params, 'validation_file_name', None)
validation_name = train_params.val_file_name
if validation_name is not None:
val_bytes = await self.minio_repository.download_file(
object_name=validation_name,
@@ -179,34 +209,35 @@ class Training(SientiaMonitoring):
wrapper = await self.plugin_store.get_model(
model_name=train_params.model_name,
force_download=False,
opt_params={},
model_kwargs={},
data_model_kwargs={},
metadata=metadata,
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
)
train_df = pd.concat([train_result.x_train, train_result.y_train], axis=1)
val_df = pd.concat([train_result.x_test, train_result.y_test], axis=1)
train_data = train_result.train_data
val_data = train_result.val_data
wrapper.train(
train_data=train_df,
val_data=val_df,
train_data=train_data,
val_data=val_data,
target=train_params.target_variable,
)
# Generate predictions using the trained wrapper
transformed_train, _ = wrapper.transform(train_df)
transformed_val, _ = wrapper.transform(val_df)
transformed_train, _ = wrapper.transform(train_data)
transformed_val, _ = wrapper.transform(val_data)
y_train_pred_df, _ = wrapper.predict({}, transformed_train)
y_val_pred_df, _ = wrapper.predict({}, transformed_val)
# Use the first column of the prediction DataFrame as the target prediction
train_result.y_train_pred = y_train_pred_df.iloc[:, 0]
train_result.y_pred = y_val_pred_df.iloc[:, 0]
y_train_pred_df.sort_index(inplace=True, ascending=False)
y_val_pred_df.sort_index(inplace=True, ascending=False)
train_result.y_train_pred = y_train_pred_df
train_result.y_pred = y_val_pred_df
train_result = self.data_manager_repository.compute_regression_metrics(
train_params,
train_result,
)
@@ -215,7 +246,7 @@ class Training(SientiaMonitoring):
async with self.mlflow_repository.start_run(
model_name=train_params.model_name,
run_name=None,
experiment_name=train_params.experiment_name,
experiment_name=f'{train_params.model_name}_experiment',
tags=None,
metadata=metadata,
) as run_info:
@@ -224,7 +255,8 @@ class Training(SientiaMonitoring):
model_saved = True
return {
'run_name': run_info.run_name or run_info.run_id,
'run_name': run_info.run_name,
'run_id': run_info.run_id,
}
except Exception as e: # noqa: BLE001
metrics_status = 'error'
@@ -237,8 +269,8 @@ class Training(SientiaMonitoring):
trace = traceback.format_exc()
self.send_notification(
metadata=metadata,
await self.send_notification_async(
metadata=metadata or {},
notification_id='TRAIN_MODEL_ERROR',
message=error_msg,
block='train_model',
@@ -250,13 +282,6 @@ class Training(SientiaMonitoring):
model_trained=model_trained,
model_saved=model_saved,
) from e
finally:
await self._emit_metrics(
metadata=metadata,
metrics_status=metrics_status,
activity_name='train_model',
emit_workflow_metric=(metrics_status == 'error'),
)
@activity.defn(name='cleanup_resources')
async def cleanup_resources(self, input_data: dict[str, Any]) -> None:
@@ -284,8 +309,6 @@ class Training(SientiaMonitoring):
metadata=metadata,
)
except Exception as e: # noqa: BLE001
metrics_status = 'error'
error_msg = (
'Error cleaning up resources - '
f'File: {bucket_name}/{file_name}, Error: {str(e)}'
@@ -303,44 +326,3 @@ class Training(SientiaMonitoring):
)
raise
finally:
await self._emit_metrics(
metadata=metadata,
metrics_status=metrics_status,
activity_name='cleanup_resources',
emit_workflow_metric=True,
)
async def _emit_metrics(
self,
metadata: dict[str, Any],
metrics_status: str,
activity_name: str,
emit_workflow_metric: bool,
) -> None:
"""
Emit workflow and activity execution metrics.
Args:
metadata: Activity metadata containing pod_id and workflow_name
metrics_status: Execution status ('success' or 'error')
activity_name: Name of the activity being executed
"""
if emit_workflow_metric:
await self.emit_metric(
metric_object=WORKFLOW_EXECUTION_TOTAL,
tags={
'pod_id': metadata.get('pod_id'),
'workflow_name': metadata.get('workflow_name'),
'status': metrics_status,
},
)
await self.emit_metric(
metric_object=ACTIVITY_EXECUTION_TOTAL,
tags={
'pod_id': metadata.get('pod_id'),
'activity_name': activity_name,
'status': metrics_status,
},
)