SIENTIAPDE-1231
SIENTIAPDE-1222 Enhance MLFlow and MLFlowRepository with model configuration support - Introduced `model_config` parameter in MLFlow methods to streamline model handling and configuration management. - Updated `retrain_model`, `transform`, and `predict` methods to accept `model_config` and `metadata` for improved flexibility and logging. - Added `detect_and_parse_datetime_index` method to handle datetime index parsing with enhanced error handling and logging. - Refactored model experiment creation to include transformation and prediction flavors, along with compression options. - Improved documentation and type hints across methods for better clarity and usability.
This commit is contained in:
@@ -225,6 +225,7 @@ class MLFlow(BaseActivity):
|
|||||||
metadata = input_data['metadata']
|
metadata = input_data['metadata']
|
||||||
data = DataFrame(input_data['data'])
|
data = DataFrame(input_data['data'])
|
||||||
model_name = input_data['model_name']
|
model_name = input_data['model_name']
|
||||||
|
model_config = input_data.get('model_config', {})
|
||||||
|
|
||||||
self.info(f'Retraining model {model_name}...', metadata)
|
self.info(f'Retraining model {model_name}...', metadata)
|
||||||
|
|
||||||
@@ -245,7 +246,8 @@ class MLFlow(BaseActivity):
|
|||||||
try:
|
try:
|
||||||
retrain_output, experiment = self.model_monitoring_repository.retrain_model(
|
retrain_output, experiment = self.model_monitoring_repository.retrain_model(
|
||||||
data=data,
|
data=data,
|
||||||
model_name=model_name
|
model_name=model_name,
|
||||||
|
model_config=model_config
|
||||||
)
|
)
|
||||||
|
|
||||||
return {
|
return {
|
||||||
|
|||||||
@@ -1,14 +1,17 @@
|
|||||||
"""
|
"""
|
||||||
Model Monitoring Repository
|
MLFlow Repository
|
||||||
|
|
||||||
This module contains the ModelMonitoringRepository class,
|
This module contains the MLFlowRepository class,
|
||||||
which is responsible for handling the communication with the Model Monitoring API.
|
which is responsible for handling the communication with MLFlow tracking server.
|
||||||
|
|
||||||
It includes the methods that are used to answer ModelMonitoringService
|
It includes methods for model management, caching, retraining, and serving operations
|
||||||
requests using the Model Monitoring API functions.
|
using MLFlow's tracking and model registry capabilities.
|
||||||
|
|
||||||
By Monitoring we mean the evaluation of the performance of models, the generation of reports.
|
|
||||||
|
|
||||||
|
The repository provides comprehensive functionality for:
|
||||||
|
- Model loading and caching with retention policies
|
||||||
|
- Data transformation and prediction operations
|
||||||
|
- Model retraining workflows
|
||||||
|
- Production model updates and versioning
|
||||||
"""
|
"""
|
||||||
from datetime import datetime, timedelta
|
from datetime import datetime, timedelta
|
||||||
import traceback
|
import traceback
|
||||||
@@ -21,6 +24,8 @@ import lzma
|
|||||||
import gzip
|
import gzip
|
||||||
import pickle
|
import pickle
|
||||||
|
|
||||||
|
from sientia_do.temporal.constants import DATETIME_FORMAT_WITH_TZ
|
||||||
|
|
||||||
ARTIFACTS_PATH = "./tmp/artifacts"
|
ARTIFACTS_PATH = "./tmp/artifacts"
|
||||||
|
|
||||||
|
|
||||||
@@ -36,6 +41,7 @@ class MLFlowRepository():
|
|||||||
self.client = mlflow.tracking.MlflowClient()
|
self.client = mlflow.tracking.MlflowClient()
|
||||||
|
|
||||||
self.model_cache = {}
|
self.model_cache = {}
|
||||||
|
self.logger = logger
|
||||||
|
|
||||||
"""
|
"""
|
||||||
Functions related to get model registry parameters
|
Functions related to get model registry parameters
|
||||||
@@ -47,6 +53,7 @@ class MLFlowRepository():
|
|||||||
|
|
||||||
Args:
|
Args:
|
||||||
run_id (str): The run_id of the model.
|
run_id (str): The run_id of the model.
|
||||||
|
prediction (bool): Whether to get prediction model URI (default: True)
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
str: The model URI.
|
str: The model URI.
|
||||||
@@ -80,7 +87,16 @@ class MLFlowRepository():
|
|||||||
run_id = latest_versions[0].source.split("/")
|
run_id = latest_versions[0].source.split("/")
|
||||||
return run_id[2]
|
return run_id[2]
|
||||||
|
|
||||||
def get_experiment_by_run_id(self, run_id: str) -> dict:
|
def get_experiment_by_run_id(self, run_id: str) -> str:
|
||||||
|
"""
|
||||||
|
Get experiment name by run ID.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
run_id (str): The MLFlow run ID
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
str: The experiment name
|
||||||
|
"""
|
||||||
# Get the run information using the run_id
|
# Get the run information using the run_id
|
||||||
run = mlflow.get_run(run_id)
|
run = mlflow.get_run(run_id)
|
||||||
|
|
||||||
@@ -183,6 +199,13 @@ class MLFlowRepository():
|
|||||||
def dowload_artifacts(self, model_name: str, artifact_path: str = "data_model") -> str:
|
def dowload_artifacts(self, model_name: str, artifact_path: str = "data_model") -> str:
|
||||||
"""
|
"""
|
||||||
Downloads artifacts from a specific MLFlow run.
|
Downloads artifacts from a specific MLFlow run.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
model_name (str): Name of the model
|
||||||
|
artifact_path (str): Path to the artifact within the run
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
str: Path to the downloaded artifacts
|
||||||
"""
|
"""
|
||||||
run_id = self.get_model_run_id(
|
run_id = self.get_model_run_id(
|
||||||
model_name=model_name, stage="Production"
|
model_name=model_name, stage="Production"
|
||||||
@@ -202,10 +225,15 @@ class MLFlowRepository():
|
|||||||
artifact_path: str | None = None):
|
artifact_path: str | None = None):
|
||||||
"""
|
"""
|
||||||
Downloads a predictive model from the MLflow Model Registry.
|
Downloads a predictive model from the MLflow Model Registry.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
model_name (str): The name of the model to download from the registry.
|
model_name (str): The name of the model to download from the registry.
|
||||||
|
flavor (str): Model flavor ('pyfunc', 'sklearn', 'pytorch')
|
||||||
|
artifact_path (str | None): Path to compressed artifacts if model is compressed
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
mlflow.pyfunc.PyFuncModel: The loaded predictive model.
|
mlflow.pyfunc.PyFuncModel: The loaded predictive model.
|
||||||
|
|
||||||
Notes:
|
Notes:
|
||||||
- The model is fetched from the "production" stage of the MLflow Model Registry.
|
- The model is fetched from the "production" stage of the MLflow Model Registry.
|
||||||
- Warnings during the model loading process are suppressed.
|
- Warnings during the model loading process are suppressed.
|
||||||
@@ -236,13 +264,15 @@ class MLFlowRepository():
|
|||||||
def load_transform_model(self, model_name: str, flavor: str,
|
def load_transform_model(self, model_name: str, flavor: str,
|
||||||
artifact_path: str | None = None):
|
artifact_path: str | None = None):
|
||||||
"""
|
"""
|
||||||
Downloads the latest production version of a specified model.
|
Downloads the latest production version of a specified transformation model.
|
||||||
|
|
||||||
This method retrieves the latest production model run ID for the given
|
This method retrieves the latest production model run ID for the given
|
||||||
model name, constructs the model URI, and loads the model using MLflow.
|
model name, constructs the model URI, and loads the model using MLflow.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
model_name (str): The name of the model to download.
|
model_name (str): The name of the model to download.
|
||||||
|
flavor (str): Model flavor ('sklearn', 'pyfunc', 'pytorch')
|
||||||
|
artifact_path (str | None): Path to compressed artifacts if model is compressed
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Any: The loaded model object, as returned by `mlflow.sklearn.load_model`.
|
Any: The loaded model object, as returned by `mlflow.sklearn.load_model`.
|
||||||
@@ -284,7 +314,8 @@ class MLFlowRepository():
|
|||||||
Load model from pickle file trying different compression methods.
|
Load model from pickle file trying different compression methods.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
pickle_path (str): Path to the pickle file
|
artifact_path (str): Path to the artifact directory
|
||||||
|
type (str): Type of model ('transformer' or 'prediction')
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Any: Loaded model object
|
Any: Loaded model object
|
||||||
@@ -346,37 +377,80 @@ class MLFlowRepository():
|
|||||||
Returns:
|
Returns:
|
||||||
dict: Model configuration with model and artifact paths
|
dict: Model configuration with model and artifact paths
|
||||||
"""
|
"""
|
||||||
if model_type == "predict":
|
|
||||||
|
if model_type not in ["predict", "transform"]:
|
||||||
|
raise ValueError(
|
||||||
|
"Invalid model_type. Use 'predict' or 'transform'.")
|
||||||
|
|
||||||
if compressed:
|
if compressed:
|
||||||
|
target = "prediction_model" if model_type == "predict" else "data_model"
|
||||||
|
|
||||||
artifact_path = self.dowload_artifacts(
|
artifact_path = self.dowload_artifacts(
|
||||||
model_name, "prediction_model")
|
model_name, target)
|
||||||
else:
|
else:
|
||||||
artifact_path = None
|
artifact_path = None
|
||||||
|
|
||||||
|
if model_type == "predict":
|
||||||
model = self.load_predict_model(model_name, flavor, artifact_path)
|
model = self.load_predict_model(model_name, flavor, artifact_path)
|
||||||
|
|
||||||
elif model_type == "transform":
|
elif model_type == "transform":
|
||||||
if compressed:
|
|
||||||
self.logger.info(
|
|
||||||
f"Model {model_name} is compressed, downloading artifacts")
|
|
||||||
|
|
||||||
artifact_path = self.dowload_artifacts(
|
|
||||||
model_name, "data_model")
|
|
||||||
|
|
||||||
self.logger.info(
|
|
||||||
f"Artifacts downloaded at path {artifact_path}")
|
|
||||||
else:
|
|
||||||
artifact_path = None
|
|
||||||
model = self.load_transform_model(
|
model = self.load_transform_model(
|
||||||
model_name, flavor, artifact_path)
|
model_name, flavor, artifact_path)
|
||||||
else:
|
|
||||||
raise ValueError(
|
|
||||||
"Invalid model_type. Use 'predict' or 'transform'.")
|
|
||||||
|
|
||||||
return {
|
return {
|
||||||
"model": model,
|
"model": model,
|
||||||
"artifact_path": artifact_path
|
"artifact_path": artifact_path
|
||||||
}
|
}
|
||||||
|
|
||||||
|
"""
|
||||||
|
Functions related to data format
|
||||||
|
"""
|
||||||
|
|
||||||
|
def detect_and_parse_datetime_index(self, data: pd.DataFrame, metadata: dict) -> pd.DataFrame:
|
||||||
|
"""
|
||||||
|
Detect and parse datetime index from data. index must be a timestamp like column.
|
||||||
|
This function must detect the timestamp type (pandas Timestamp or datetime) and convert it to DATETIME_FORMAT_WITH_TZ.
|
||||||
|
If the index is a string, must be in format DATETIME_FORMAT_WITH_TZ.
|
||||||
|
If another type or format, must raise an error.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
data (pd.DataFrame): DataFrame with timestamp index
|
||||||
|
metadata (dict): Metadata for logging
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
pd.DataFrame: DataFrame with converted datetime index
|
||||||
|
"""
|
||||||
|
index = data.index
|
||||||
|
|
||||||
|
# Get type of first element of index
|
||||||
|
index_type = type(index[0])
|
||||||
|
|
||||||
|
self.logger.custom_info(f"Index type: {index_type}", metadata)
|
||||||
|
|
||||||
|
message = f"Index must be all timestamp like column. Valid formats are: pandas Timestamp, datetime, string in format {DATETIME_FORMAT_WITH_TZ}"
|
||||||
|
|
||||||
|
# Check if all in index are of the same type
|
||||||
|
if not all(isinstance(i, index_type) for i in index):
|
||||||
|
raise ValueError(
|
||||||
|
f"{message}")
|
||||||
|
|
||||||
|
# Check type and converts to DATETIME_FORMAT_WITH_TZ
|
||||||
|
if index_type == str:
|
||||||
|
# Validate format of string and return error if not valid
|
||||||
|
try:
|
||||||
|
pd.to_datetime(data.index, format=DATETIME_FORMAT_WITH_TZ)
|
||||||
|
except ValueError:
|
||||||
|
raise ValueError(
|
||||||
|
f"{message}")
|
||||||
|
|
||||||
|
elif index_type == datetime or index_type == pd.Timestamp:
|
||||||
|
data.index = data.index.strftime(DATETIME_FORMAT_WITH_TZ)
|
||||||
|
else:
|
||||||
|
raise ValueError(
|
||||||
|
f"{message}")
|
||||||
|
|
||||||
|
return data
|
||||||
|
|
||||||
"""
|
"""
|
||||||
Functions related to cache management of models
|
Functions related to cache management of models
|
||||||
"""
|
"""
|
||||||
@@ -445,7 +519,7 @@ class MLFlowRepository():
|
|||||||
else:
|
else:
|
||||||
return cache['target']
|
return cache['target']
|
||||||
|
|
||||||
def handle_outdated_model(self, model_name: str, model_key: str) -> dict:
|
def handle_outdated_model(self, model_name: str, model_key: str) -> None:
|
||||||
"""
|
"""
|
||||||
Clean up outdated cached model and its artifacts.
|
Clean up outdated cached model and its artifacts.
|
||||||
|
|
||||||
@@ -454,7 +528,7 @@ class MLFlowRepository():
|
|||||||
model_key (str): Cache key for the model
|
model_key (str): Cache key for the model
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
dict: Empty dictionary (cleanup operation)
|
None
|
||||||
"""
|
"""
|
||||||
if self.logger:
|
if self.logger:
|
||||||
self.logger.debug(
|
self.logger.debug(
|
||||||
@@ -596,7 +670,9 @@ class MLFlowRepository():
|
|||||||
Functions related to model retraining
|
Functions related to model retraining
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def create_model_experiment(self, model_name: str, data: pd.DataFrame) -> tuple:
|
def create_model_experiment(self, model_name: str, data: pd.DataFrame,
|
||||||
|
transform_flavor: str = 'sklearn', predict_flavor: str = 'pyfunc',
|
||||||
|
compressed: bool = False) -> tuple:
|
||||||
"""
|
"""
|
||||||
Create a new MLFlow experiment for model retraining.
|
Create a new MLFlow experiment for model retraining.
|
||||||
|
|
||||||
@@ -610,6 +686,9 @@ class MLFlowRepository():
|
|||||||
Args:
|
Args:
|
||||||
model_name (str): Name of the MLFlow model to retrain
|
model_name (str): Name of the MLFlow model to retrain
|
||||||
data (pd.DataFrame): Training data for model retraining
|
data (pd.DataFrame): Training data for model retraining
|
||||||
|
transform_flavor (str): Flavor for transformation model
|
||||||
|
predict_flavor (str): Flavor for prediction model
|
||||||
|
compressed (bool): Whether models are compressed
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
tuple: (prediction_model, data_model, experiment)
|
tuple: (prediction_model, data_model, experiment)
|
||||||
@@ -617,18 +696,15 @@ class MLFlowRepository():
|
|||||||
- data_model: Fitted transformation model
|
- data_model: Fitted transformation model
|
||||||
- experiment: MLFlow experiment name
|
- experiment: MLFlow experiment name
|
||||||
"""
|
"""
|
||||||
# load predictor model
|
|
||||||
predictor_uri = f"models:/{model_name}/production"
|
|
||||||
# load transform model
|
|
||||||
latest_production_id = self.get_model_run_id(
|
latest_production_id = self.get_model_run_id(
|
||||||
model_name, stage="Production"
|
model_name, stage="Production"
|
||||||
)
|
)
|
||||||
transform_uri = self.get_model_uri(
|
data_model = self.download_model(
|
||||||
latest_production_id, prediction=False
|
model_name, "transform", transform_flavor, compressed
|
||||||
|
)
|
||||||
|
prediction_model = self.download_model(
|
||||||
|
model_name, "predict", predict_flavor, compressed
|
||||||
)
|
)
|
||||||
# load
|
|
||||||
data_model = mlflow.sklearn.load_model(transform_uri)
|
|
||||||
prediction_model = mlflow.sklearn.load_model(predictor_uri)
|
|
||||||
data_model = data_model.fit(data)
|
data_model = data_model.fit(data)
|
||||||
treated_data = data_model.predict(data)
|
treated_data = data_model.predict(data)
|
||||||
|
|
||||||
@@ -768,7 +844,8 @@ class MLFlowRepository():
|
|||||||
Functions that provide the interface to model operations
|
Functions that provide the interface to model operations
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def transform(self, model_name: str, data: pd.DataFrame, model_retention: int):
|
def transform(self, model_name: str, data: pd.DataFrame, model_retention: int,
|
||||||
|
model_config: dict, metadata: dict):
|
||||||
"""
|
"""
|
||||||
Transform data using a cached transformation model.
|
Transform data using a cached transformation model.
|
||||||
|
|
||||||
@@ -787,6 +864,8 @@ class MLFlowRepository():
|
|||||||
model_name (str): The name of the MLFlow model to use for transformation.
|
model_name (str): The name of the MLFlow model to use for transformation.
|
||||||
data (pd.DataFrame): The input data to be transformed by the model.
|
data (pd.DataFrame): The input data to be transformed by the model.
|
||||||
model_retention (int): Cache retention time in minutes (0 = no caching).
|
model_retention (int): Cache retention time in minutes (0 = no caching).
|
||||||
|
model_config (dict): Model configuration parameters
|
||||||
|
metadata (dict): Metadata for logging
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
dict: Response dictionary containing:
|
dict: Response dictionary containing:
|
||||||
@@ -800,13 +879,31 @@ class MLFlowRepository():
|
|||||||
Exception: Any exception during model loading or transformation is caught
|
Exception: Any exception during model loading or transformation is caught
|
||||||
and returned in the response structure rather than propagated.
|
and returned in the response structure rather than propagated.
|
||||||
"""
|
"""
|
||||||
|
self.logger.custom_debug(
|
||||||
|
f"Data received for model transformation: {data.to_csv()}", metadata)
|
||||||
|
|
||||||
|
model_retention = model_config.get('retention_minutes', 0)
|
||||||
|
flavor = model_config.get('transform_flavor', 'sklearn')
|
||||||
|
compressed = model_config.get('is_compressed', False)
|
||||||
|
retention_target = model_config.get('retention_target', 'model')
|
||||||
|
transform_keyword = model_config.get(
|
||||||
|
'transform_function_keyword', 'predict')
|
||||||
|
|
||||||
try:
|
try:
|
||||||
|
transformed_data = self.get_cached_transform(
|
||||||
|
model_name, data, model_retention, flavor,
|
||||||
|
compressed, retention_target, transform_keyword
|
||||||
|
)
|
||||||
|
|
||||||
|
self.logger.custom_debug(
|
||||||
|
f"Data received from model transformation: {transformed_data.to_csv()}", metadata)
|
||||||
|
|
||||||
|
transformed_data = self.detect_and_parse_datetime_index(
|
||||||
|
transformed_data, metadata)
|
||||||
|
|
||||||
return {
|
return {
|
||||||
'success': True,
|
'success': True,
|
||||||
'content': self.get_cached_transform(
|
'content': transformed_data.to_dict()
|
||||||
model_name, data, model_retention).to_dict()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
@@ -818,7 +915,8 @@ class MLFlowRepository():
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
def predict(self, model_name: str, data: pd.DataFrame, model_retention: int):
|
def predict(self, model_name: str, data: pd.DataFrame, model_retention: int,
|
||||||
|
model_config: dict, metadata: dict):
|
||||||
"""
|
"""
|
||||||
Generate predictions using a cached prediction model.
|
Generate predictions using a cached prediction model.
|
||||||
|
|
||||||
@@ -841,6 +939,8 @@ class MLFlowRepository():
|
|||||||
model_name (str): The name of the MLFlow model to use for prediction.
|
model_name (str): The name of the MLFlow model to use for prediction.
|
||||||
data (pd.DataFrame): The input data to make predictions on.
|
data (pd.DataFrame): The input data to make predictions on.
|
||||||
model_retention (int): Cache retention time in minutes (0 = no caching).
|
model_retention (int): Cache retention time in minutes (0 = no caching).
|
||||||
|
model_config (dict): Model configuration parameters
|
||||||
|
metadata (dict): Metadata for logging
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
dict: Response dictionary containing:
|
dict: Response dictionary containing:
|
||||||
@@ -856,15 +956,24 @@ class MLFlowRepository():
|
|||||||
Exception: Any exception during model loading or prediction is caught
|
Exception: Any exception during model loading or prediction is caught
|
||||||
and returned in the response structure rather than propagated.
|
and returned in the response structure rather than propagated.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
model_retention = model_config.get('retention_minutes', 0)
|
||||||
|
flavor = model_config.get('predict_flavor', 'pyfunc')
|
||||||
|
|
||||||
try:
|
try:
|
||||||
|
|
||||||
input_index = data.index
|
input_index = data.index
|
||||||
start_time = datetime.now()
|
start_time = datetime.now()
|
||||||
|
|
||||||
|
self.logger.custom_debug(
|
||||||
|
f"Data received for model prediction: {data.to_csv()}", metadata)
|
||||||
data = self.get_cached_predict(
|
data = self.get_cached_predict(
|
||||||
model_name, data, model_retention)
|
model_name, data, model_retention, flavor)
|
||||||
|
|
||||||
end_time = datetime.now()
|
end_time = datetime.now()
|
||||||
data = pd.DataFrame(data, columns=['prediction'])
|
data = pd.DataFrame(data, columns=['prediction'])
|
||||||
|
self.logger.custom_debug(
|
||||||
|
f"Data received from model prediction: {data.to_csv()}", metadata)
|
||||||
data.index = input_index
|
data.index = input_index
|
||||||
data['response_time'] = (end_time - start_time).total_seconds()
|
data['response_time'] = (end_time - start_time).total_seconds()
|
||||||
|
|
||||||
@@ -882,7 +991,8 @@ class MLFlowRepository():
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
def retrain_model(self, data: pd.DataFrame, model_name: str) -> tuple:
|
def retrain_model(self, data: pd.DataFrame, model_name: str,
|
||||||
|
model_config: dict) -> tuple:
|
||||||
"""
|
"""
|
||||||
Orchestrate the complete model retraining workflow.
|
Orchestrate the complete model retraining workflow.
|
||||||
|
|
||||||
@@ -913,6 +1023,7 @@ class MLFlowRepository():
|
|||||||
prediction models, including target variable.
|
prediction models, including target variable.
|
||||||
model_name (str): Name of the MLFlow model to retrain. Must exist
|
model_name (str): Name of the MLFlow model to retrain. Must exist
|
||||||
in the MLFlow Model Registry in Production stage.
|
in the MLFlow Model Registry in Production stage.
|
||||||
|
model_config (dict): Model configuration parameters
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
tuple: Retraining operation results containing:
|
tuple: Retraining operation results containing:
|
||||||
@@ -924,8 +1035,12 @@ class MLFlowRepository():
|
|||||||
ValueError: If experiment cannot be created or models cannot be loaded
|
ValueError: If experiment cannot be created or models cannot be loaded
|
||||||
Exception: Any other exception during the retraining process
|
Exception: Any other exception during the retraining process
|
||||||
"""
|
"""
|
||||||
|
transform_flavor = model_config.get('transform_flavor', 'sklearn')
|
||||||
|
predict_flavor = model_config.get('predict_flavor', 'pyfunc')
|
||||||
|
compressed = model_config.get('is_compressed', False)
|
||||||
|
|
||||||
prediction_model, data_model, experiment = self.create_model_experiment(
|
prediction_model, data_model, experiment = self.create_model_experiment(
|
||||||
model_name, data)
|
model_name, data, transform_flavor, predict_flavor, compressed)
|
||||||
retrain_result = self.perform_model_retrain(
|
retrain_result = self.perform_model_retrain(
|
||||||
prediction_model, data_model, experiment, model_name, data)
|
prediction_model, data_model, experiment, model_name, data)
|
||||||
return retrain_result
|
return retrain_result
|
||||||
|
|||||||
Reference in New Issue
Block a user