diff --git a/laborious/utils/repository/minio_repository.py b/laborious/utils/repository/minio_repository.py index df20a06..b38af69 100644 --- a/laborious/utils/repository/minio_repository.py +++ b/laborious/utils/repository/minio_repository.py @@ -56,7 +56,7 @@ class MinioRepository: def close(self): 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. """ @@ -64,13 +64,10 @@ class MinioRepository: try: self.logger.custom_info(f"Checking if bucket '{self.minio_bucket}' exists", metadata) self.s3_client.head_bucket(Bucket=self.minio_bucket) - return True except ClientError: self.logger.custom_info(f"Creating bucket '{self.minio_bucket}'", metadata) self.s3_client.create_bucket(Bucket=self.minio_bucket) - return True - def store_dataframe_as_parquet( self, dataframe: DataFrame, uri: str, object_name: str, metadata: dict[str, Any] ): diff --git a/laborious/utils/repository/model_repository.py b/laborious/utils/repository/model_repository.py index 3e786c7..f9e783b 100644 --- a/laborious/utils/repository/model_repository.py +++ b/laborious/utils/repository/model_repository.py @@ -33,6 +33,7 @@ ARTIFACTS_PATH = './tmp/artifacts' TRANSFORMED_COMPRESSED_PATH = 'artifacts/training_transformer.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): gc.collect() @@ -222,7 +223,7 @@ class MLFlowRepository: elif flavor == 'pytorch': model = mlflow.pytorch.load_model(model_uri) else: - raise ValueError("Invalid flavor. Use 'sklearn' or 'pyfunc' or 'pytorch'.") + raise ValueError(INVALID_FLAVOR_MESSAGE) return model @@ -257,7 +258,7 @@ class MLFlowRepository: elif flavor == 'pytorch': model = mlflow.pytorch.load_model(model_uri) else: - raise ValueError("Invalid flavor. Use 'sklearn' or 'pyfunc' or 'pytorch'.") + raise ValueError(INVALID_FLAVOR_MESSAGE) return model def download_model( @@ -628,7 +629,7 @@ class MLFlowRepository: # Aligns data with treated data indexes to get target variable aligned_data = data.loc[treated_data.index] 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 ) else: @@ -673,7 +674,7 @@ class MLFlowRepository: elif flavor == 'pytorch': mlflow.pytorch.log_model(model, model_type) else: - raise ValueError("Invalid flavor. Use 'sklearn' or 'pyfunc' or 'pytorch'.") + raise ValueError(INVALID_FLAVOR_MESSAGE) def create_new_experiment( self, @@ -980,7 +981,7 @@ class MLFlowRepository: # predict_data.to_csv( # f"tmp/predicted_data_{model_name}.csv", index=True) - predict_data.columns = ['prediction'] + predict_data.columns = pd.Index(['prediction']) else: predict_data = pd.DataFrame(predict_data, columns=['prediction'])