diff --git a/laborious/utils/repository/model_repository.py b/laborious/utils/repository/model_repository.py index 337eda2..581e57f 100644 --- a/laborious/utils/repository/model_repository.py +++ b/laborious/utils/repository/model_repository.py @@ -36,6 +36,7 @@ from sientia_do.observability.sientia_monitoring import SientiaMonitoring from sientia_do.temporal.constants import DATETIME_FORMAT_WITH_TZ from laborious import metrics +from laborious.utils.dataframe_debug import build_dataframe_debug_message ARTIFACTS_PATH = './tmp/artifacts' TRANSFORMED_COMPRESSED_PATH = 'artifacts/training_transformer.pkl' @@ -64,6 +65,8 @@ def force_memory_release(logger: Logger): class MLFlowRepository(SientiaMonitoring): + _MAX_DEBUG_DATAFRAME_ROWS = 100 + def __init__( self, host: str, @@ -94,6 +97,24 @@ class MLFlowRepository(SientiaMonitoring): self._cache_lock = threading.RLock() self.logger = logger + def _debug_dataframe(self, message: str, data: Any, metadata: dict[str, Any]) -> None: + """ + Log dataframe content only when row count is below the configured threshold + + Args: + - message (str): Base log message to identify the dataframe in logs + - data (Any): Dataframe-like object expected to expose shape and to_csv + - metadata (dict[str, Any]): Metadata for contextual logging + """ + self.debug( + build_dataframe_debug_message( + message=message, + data=data, + max_rows=self._MAX_DEBUG_DATAFRAME_ROWS, + ), + metadata, + ) + """ Functions related to get model registry parameters """ @@ -438,17 +459,28 @@ class MLFlowRepository(SientiaMonitoring): raise ValueError("Invalid model_type. Use 'predict' or 'transform'.") artifact_path = None - if model_type == 'predict': - model = await self.load_predict_model(model_name, metadata, flavor) - else: - model = await self.load_transform_model(model_name, metadata, flavor) if load_wrapper: self.info(f'Loading wrapper for {model_type} model {model_name} with flavor {flavor}') - model = model._model_impl.python_model + target = 'prediction_model' if model_type == 'predict' else 'data_model' + + artifact_path = await self.dowload_artifacts(model_name, metadata, target) + + self.info( + f'Model with type {model_type} and name {model_name} is compressed, loading from {artifact_path}' + ) + + raw_model = mlflow.pyfunc.load_model(artifact_path) + model = raw_model._model_impl.python_model self.debug(f"Model wrapper loaded: {model.__class__.__name__}:{model.__dict__}", metadata) + else: + if model_type == 'predict': + model = await self.load_predict_model(model_name, metadata, flavor) + + else: + model = await self.load_transform_model(model_name, metadata, flavor) return model, artifact_path @@ -1125,7 +1157,7 @@ class MLFlowRepository(SientiaMonitoring): async def transform( self, model_name: str, data: pd.DataFrame, model_config: dict, metadata: dict - ): + ) -> dict[str, Any]: """ Transform data using a cached transformation model. @@ -1159,7 +1191,7 @@ class MLFlowRepository(SientiaMonitoring): and returned in the response structure rather than propagated. """ - self.debug(f'Data received for model transformation: {data.to_csv()}', metadata) + self._debug_dataframe('Data received for model transformation:', data, metadata) # data.to_csv( # f"tmp/data_{model_name}.csv", index=True) @@ -1177,9 +1209,8 @@ class MLFlowRepository(SientiaMonitoring): metadata=metadata, ) - self.debug( - f'Data received from model transformation: {transformed_data.head(5).to_csv()}', - metadata, + self._debug_dataframe( + 'Data received from model transformation:', transformed_data, metadata ) # transformed_data.to_csv( @@ -1187,7 +1218,7 @@ class MLFlowRepository(SientiaMonitoring): transformed_data = self.detect_and_parse_datetime_index(transformed_data, metadata) - return {'success': True, 'content': transformed_data.to_dict()} + return {'success': True, 'content': transformed_data} except Exception as e: return { @@ -1245,9 +1276,7 @@ class MLFlowRepository(SientiaMonitoring): input_index = data.index start_time = datetime.now() - self.debug( - f'Data received for model prediction: {data.to_dict(orient="records")}', metadata - ) + self._debug_dataframe('Data received for model prediction:', data, metadata) # data.to_csv( # f"tmp/treated_data_{model_name}.csv", index=True) @@ -1264,16 +1293,17 @@ class MLFlowRepository(SientiaMonitoring): end_time = datetime.now() if isinstance(predict_data, pd.DataFrame): - self.debug( - f'Data received from model prediction: {predict_data.to_dict(orient="records")}', - metadata, - ) + self._debug_dataframe('Data received from model prediction:', predict_data, metadata) # predict_data.to_csv( # f"tmp/predicted_data_{model_name}.csv", index=True) predict_data.columns = pd.Index(['prediction']) else: + self.debug( + f'Data received from model prediction (not a DataFrame): {predict_data}', + metadata, + ) predict_data = pd.DataFrame(predict_data, columns=['prediction']) # predict_data.to_csv( # f"tmp/predicted_data_{model_name}.csv", index=True) @@ -1281,7 +1311,7 @@ class MLFlowRepository(SientiaMonitoring): predict_data.index = input_index predict_data['response_time'] = (end_time - start_time).total_seconds() - return {'success': True, 'content': predict_data.to_dict()} + return {'success': True, 'content': predict_data} except Exception as e: return { @@ -1335,7 +1365,7 @@ class MLFlowRepository(SientiaMonitoring): """ self.info(f'Starting model retraining workflow for {model_name}', metadata) - self.debug(f'Data received for model retraining: {data.to_csv()}', metadata) + self._debug_dataframe('Data received for model retraining:', data, metadata) target_name = model_config.get('target', None)