diff --git a/model_manager/activities/training.py b/model_manager/activities/training.py index 12302e5..fb14f98 100644 --- a/model_manager/activities/training.py +++ b/model_manager/activities/training.py @@ -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)