feat: integrate SientiaModel wrapper and enhance logging in Training class
- Imported SientiaModel to standardize the wrapper type in the Training class. - Added conditional logging to capture training process details when a logger is provided, improving traceability during model training.
This commit is contained in:
@@ -6,6 +6,7 @@ The activity extends BaseActivity and receives pre-downloaded files
|
||||
and raises `ModelTrainingError` when training fails.
|
||||
"""
|
||||
|
||||
from sientia_model.wrappers.sientia_model import SientiaModel
|
||||
from temporalio import activity, workflow
|
||||
|
||||
with workflow.unsafe.imports_passed_through():
|
||||
@@ -228,6 +229,9 @@ class Training(SientiaMonitoring):
|
||||
metadata=metadata,
|
||||
)
|
||||
|
||||
if self.logger is not None:
|
||||
wrapper.logger = self.logger.base_logger
|
||||
|
||||
self.info(f'Training model for {train_params.model_type}', metadata)
|
||||
train_data = train_result.train_data
|
||||
val_data = train_result.val_data
|
||||
@@ -317,7 +321,7 @@ class Training(SientiaMonitoring):
|
||||
self,
|
||||
train_result: TrainModelResult,
|
||||
train_params: TrainModelParams,
|
||||
wrapper: Any,
|
||||
wrapper: SientiaModel,
|
||||
metadata: dict[str, Any] | None,
|
||||
) -> None:
|
||||
self.info(f'Generating report for {train_params.model_type}', metadata)
|
||||
|
||||
Reference in New Issue
Block a user