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