SIENTIAPDE-1231
Refactor MinioRepository and MLFlowRepository for improved functionality and error handling - Updated `ensure_bucket_exists` method in MinioRepository to return None instead of a boolean, streamlining bucket existence checks. - Replaced hardcoded error messages in MLFlowRepository with a constant for better maintainability. - Adjusted column assignment in MLFlowRepository to use pd.Index for improved clarity.
This commit is contained in:
@@ -56,7 +56,7 @@ class MinioRepository:
|
|||||||
def close(self):
|
def close(self):
|
||||||
self.s3_client.close()
|
self.s3_client.close()
|
||||||
|
|
||||||
def ensure_bucket_exists(self, metadata: dict[str, Any]) -> bool:
|
def ensure_bucket_exists(self, metadata: dict[str, Any]) -> None:
|
||||||
"""
|
"""
|
||||||
Ensure the MinIO bucket exists; create it if necessary.
|
Ensure the MinIO bucket exists; create it if necessary.
|
||||||
"""
|
"""
|
||||||
@@ -64,13 +64,10 @@ class MinioRepository:
|
|||||||
try:
|
try:
|
||||||
self.logger.custom_info(f"Checking if bucket '{self.minio_bucket}' exists", metadata)
|
self.logger.custom_info(f"Checking if bucket '{self.minio_bucket}' exists", metadata)
|
||||||
self.s3_client.head_bucket(Bucket=self.minio_bucket)
|
self.s3_client.head_bucket(Bucket=self.minio_bucket)
|
||||||
return True
|
|
||||||
except ClientError:
|
except ClientError:
|
||||||
self.logger.custom_info(f"Creating bucket '{self.minio_bucket}'", metadata)
|
self.logger.custom_info(f"Creating bucket '{self.minio_bucket}'", metadata)
|
||||||
self.s3_client.create_bucket(Bucket=self.minio_bucket)
|
self.s3_client.create_bucket(Bucket=self.minio_bucket)
|
||||||
|
|
||||||
return True
|
|
||||||
|
|
||||||
def store_dataframe_as_parquet(
|
def store_dataframe_as_parquet(
|
||||||
self, dataframe: DataFrame, uri: str, object_name: str, metadata: dict[str, Any]
|
self, dataframe: DataFrame, uri: str, object_name: str, metadata: dict[str, Any]
|
||||||
):
|
):
|
||||||
|
|||||||
@@ -33,6 +33,7 @@ ARTIFACTS_PATH = './tmp/artifacts'
|
|||||||
TRANSFORMED_COMPRESSED_PATH = 'artifacts/training_transformer.pkl'
|
TRANSFORMED_COMPRESSED_PATH = 'artifacts/training_transformer.pkl'
|
||||||
PREDICTION_COMPRESSED_PATH = 'artifacts/stacking_model.pkl'
|
PREDICTION_COMPRESSED_PATH = 'artifacts/stacking_model.pkl'
|
||||||
|
|
||||||
|
INVALID_FLAVOR_MESSAGE = "Invalid flavor. Use 'sklearn' or 'pyfunc' or 'pytorch'."
|
||||||
|
|
||||||
def force_memory_release(logger: Logger):
|
def force_memory_release(logger: Logger):
|
||||||
gc.collect()
|
gc.collect()
|
||||||
@@ -222,7 +223,7 @@ class MLFlowRepository:
|
|||||||
elif flavor == 'pytorch':
|
elif flavor == 'pytorch':
|
||||||
model = mlflow.pytorch.load_model(model_uri)
|
model = mlflow.pytorch.load_model(model_uri)
|
||||||
else:
|
else:
|
||||||
raise ValueError("Invalid flavor. Use 'sklearn' or 'pyfunc' or 'pytorch'.")
|
raise ValueError(INVALID_FLAVOR_MESSAGE)
|
||||||
|
|
||||||
return model
|
return model
|
||||||
|
|
||||||
@@ -257,7 +258,7 @@ class MLFlowRepository:
|
|||||||
elif flavor == 'pytorch':
|
elif flavor == 'pytorch':
|
||||||
model = mlflow.pytorch.load_model(model_uri)
|
model = mlflow.pytorch.load_model(model_uri)
|
||||||
else:
|
else:
|
||||||
raise ValueError("Invalid flavor. Use 'sklearn' or 'pyfunc' or 'pytorch'.")
|
raise ValueError(INVALID_FLAVOR_MESSAGE)
|
||||||
return model
|
return model
|
||||||
|
|
||||||
def download_model(
|
def download_model(
|
||||||
@@ -628,7 +629,7 @@ class MLFlowRepository:
|
|||||||
# Aligns data with treated data indexes to get target variable
|
# Aligns data with treated data indexes to get target variable
|
||||||
aligned_data = data.loc[treated_data.index]
|
aligned_data = data.loc[treated_data.index]
|
||||||
aligned_series = aligned_data[target_name]
|
aligned_series = aligned_data[target_name]
|
||||||
retrain_dataset = pd.merge(
|
retrain_dataset = pd.merge( # NOSONAR
|
||||||
treated_data, aligned_series, left_index=True, right_index=True
|
treated_data, aligned_series, left_index=True, right_index=True
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
@@ -673,7 +674,7 @@ class MLFlowRepository:
|
|||||||
elif flavor == 'pytorch':
|
elif flavor == 'pytorch':
|
||||||
mlflow.pytorch.log_model(model, model_type)
|
mlflow.pytorch.log_model(model, model_type)
|
||||||
else:
|
else:
|
||||||
raise ValueError("Invalid flavor. Use 'sklearn' or 'pyfunc' or 'pytorch'.")
|
raise ValueError(INVALID_FLAVOR_MESSAGE)
|
||||||
|
|
||||||
def create_new_experiment(
|
def create_new_experiment(
|
||||||
self,
|
self,
|
||||||
@@ -980,7 +981,7 @@ class MLFlowRepository:
|
|||||||
|
|
||||||
# predict_data.to_csv(
|
# predict_data.to_csv(
|
||||||
# f"tmp/predicted_data_{model_name}.csv", index=True)
|
# f"tmp/predicted_data_{model_name}.csv", index=True)
|
||||||
predict_data.columns = ['prediction']
|
predict_data.columns = pd.Index(['prediction'])
|
||||||
|
|
||||||
else:
|
else:
|
||||||
predict_data = pd.DataFrame(predict_data, columns=['prediction'])
|
predict_data = pd.DataFrame(predict_data, columns=['prediction'])
|
||||||
|
|||||||
Reference in New Issue
Block a user