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.
|
and raises `ModelTrainingError` when training fails.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
from sientia_model.wrappers.sientia_model import SientiaModel
|
||||||
from temporalio import activity, workflow
|
from temporalio import activity, workflow
|
||||||
|
|
||||||
with workflow.unsafe.imports_passed_through():
|
with workflow.unsafe.imports_passed_through():
|
||||||
@@ -228,6 +229,9 @@ class Training(SientiaMonitoring):
|
|||||||
metadata=metadata,
|
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)
|
self.info(f'Training model for {train_params.model_type}', metadata)
|
||||||
train_data = train_result.train_data
|
train_data = train_result.train_data
|
||||||
val_data = train_result.val_data
|
val_data = train_result.val_data
|
||||||
@@ -317,7 +321,7 @@ class Training(SientiaMonitoring):
|
|||||||
self,
|
self,
|
||||||
train_result: TrainModelResult,
|
train_result: TrainModelResult,
|
||||||
train_params: TrainModelParams,
|
train_params: TrainModelParams,
|
||||||
wrapper: Any,
|
wrapper: SientiaModel,
|
||||||
metadata: dict[str, Any] | None,
|
metadata: dict[str, Any] | None,
|
||||||
) -> None:
|
) -> None:
|
||||||
self.info(f'Generating report for {train_params.model_type}', metadata)
|
self.info(f'Generating report for {train_params.model_type}', metadata)
|
||||||
|
|||||||
Reference in New Issue
Block a user