feat: integrate PluginStore and MinIO repository into model manager activities
- Added PluginStore integration for model management. - Replaced StorageRepository with MinIORepository in Activities, Cleanup, and Training classes. - Updated training logic to handle validation files and improved data management. - Enhanced configuration for MinIO and PluginStore in connectors. - Removed deprecated model repository and storage repository files. - Updated environment variable handling for new configurations.
This commit is contained in:
@@ -6,12 +6,14 @@ with workflow.unsafe.imports_passed_through():
|
||||
from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler
|
||||
from sientia_do.observability.logger import Logger
|
||||
from sientia_do.observability.metrics_controller import MetricsController
|
||||
from sientia_do.observability.sientia_monitoring import SientiaMonitoring
|
||||
from sientia_do.repository.minio_repository import MinioRepository
|
||||
from sientia_model.model_repository.mlflow_repository import SientiaMLflowRepository
|
||||
from sientia_model.model_repository.plugin_store import PluginStore
|
||||
|
||||
from model_manager.activities.cleanup import Cleanup
|
||||
from model_manager.activities.experiment_tracking import ExperimentTracking
|
||||
from model_manager.activities.training import Training
|
||||
from model_manager.utils.repository.model_repository import ModelRepository
|
||||
from model_manager.utils.repository.storage_repository import StorageRepository
|
||||
|
||||
|
||||
class Activities(ExperimentTracking, Training, Cleanup):
|
||||
@@ -40,8 +42,10 @@ class Activities(ExperimentTracking, Training, Cleanup):
|
||||
postgres_config: dict[str, Any],
|
||||
mlflow_config: dict[str, Any],
|
||||
minio_config: dict[str, Any],
|
||||
plugin_store: PluginStore,
|
||||
logger: Logger,
|
||||
notification_handler: NotificationHandler,
|
||||
metrics_controller: MetricsController,
|
||||
):
|
||||
"""
|
||||
Initialize the Activities orchestrator with all required configurations.
|
||||
@@ -62,9 +66,7 @@ class Activities(ExperimentTracking, Training, Cleanup):
|
||||
Raises:
|
||||
Exception: If any parent class initialization fails
|
||||
"""
|
||||
metrics_controller = MetricsController(
|
||||
logger=logger,
|
||||
)
|
||||
|
||||
|
||||
ExperimentTracking.__init__(
|
||||
self,
|
||||
@@ -80,30 +82,40 @@ class Activities(ExperimentTracking, Training, Cleanup):
|
||||
metrics_controller=metrics_controller,
|
||||
)
|
||||
|
||||
self.model_repository = ModelRepository(
|
||||
url=mlflow_config['url'],
|
||||
self.mlflow_repository = SientiaMLflowRepository(
|
||||
host=mlflow_config['url'],
|
||||
username=mlflow_config['username'],
|
||||
password=mlflow_config['password'],
|
||||
logger=logger,
|
||||
notification_handler=notification_handler,
|
||||
metrics_controller=metrics_controller,
|
||||
)
|
||||
|
||||
self.storage_repository = StorageRepository(
|
||||
endpoint_url=minio_config['endpoint_url'],
|
||||
# MinIO repository used for all object storage operations
|
||||
endpoint_url = minio_config['endpoint_url']
|
||||
# MinioRepository expects the endpoint without scheme
|
||||
if endpoint_url.startswith('http://'):
|
||||
endpoint = endpoint_url.removeprefix('http://')
|
||||
elif endpoint_url.startswith('https://'):
|
||||
endpoint = endpoint_url.removeprefix('https://')
|
||||
else:
|
||||
endpoint = endpoint_url
|
||||
|
||||
self.minio_repository = MinioRepository(
|
||||
endpoint=endpoint,
|
||||
access_key=minio_config['access_key'],
|
||||
secret_key=minio_config['secret_key'],
|
||||
region=minio_config['region'],
|
||||
use_ssl=minio_config['use_ssl'],
|
||||
max_retry_attempts=minio_config['max_retry_attempts'],
|
||||
retry_mode=minio_config['retry_mode'],
|
||||
connect_timeout=minio_config['connect_timeout'],
|
||||
read_timeout=minio_config['read_timeout'],
|
||||
logger=logger,
|
||||
notification_handler=notification_handler,
|
||||
metrics_controller=metrics_controller,
|
||||
secure=minio_config['use_ssl'],
|
||||
)
|
||||
|
||||
Training.__init__(
|
||||
self,
|
||||
model_repository=self.model_repository,
|
||||
storage_repository=self.storage_repository,
|
||||
mlflow_repository=self.mlflow_repository,
|
||||
plugin_store=plugin_store,
|
||||
minio_repository=self.minio_repository,
|
||||
logger=logger,
|
||||
notification_handler=notification_handler,
|
||||
metrics_controller=metrics_controller,
|
||||
@@ -111,7 +123,7 @@ class Activities(ExperimentTracking, Training, Cleanup):
|
||||
|
||||
Cleanup.__init__(
|
||||
self,
|
||||
storage_repository=self.storage_repository,
|
||||
minio_repository=self.minio_repository,
|
||||
logger=logger,
|
||||
notification_handler=notification_handler,
|
||||
metrics_controller=metrics_controller,
|
||||
@@ -137,7 +149,7 @@ class Activities(ExperimentTracking, Training, Cleanup):
|
||||
# Logging here could cause issues if logger is already destroyed
|
||||
pass
|
||||
|
||||
async def shutdown(self):
|
||||
def shutdown(self):
|
||||
"""
|
||||
Gracefully shutdown all activities and clean up resources.
|
||||
|
||||
@@ -152,4 +164,4 @@ class Activities(ExperimentTracking, Training, Cleanup):
|
||||
"""
|
||||
ExperimentTracking.close(self)
|
||||
self.info('Postgres client closed')
|
||||
self.storage_repository.close()
|
||||
SientiaMonitoring.shutdown(self)
|
||||
|
||||
@@ -21,9 +21,9 @@ with workflow.unsafe.imports_passed_through():
|
||||
from sientia_do.observability.logger import Logger
|
||||
from sientia_do.observability.metrics_controller import MetricsController
|
||||
from sientia_do.observability.sientia_monitoring import SientiaMonitoring
|
||||
from sientia_do.repository.minio_repository import MinioRepository
|
||||
|
||||
from model_manager.metrics import ACTIVITY_EXECUTION_TOTAL, WORKFLOW_EXECUTION_TOTAL
|
||||
from model_manager.utils.repository.storage_repository import StorageRepository
|
||||
|
||||
RETENTION_HOURS = int(os.getenv('CLEANUP_RETENTION_HOURS', '24'))
|
||||
DRY_RUN = os.getenv('CLEANUP_DRY_RUN', 'false').lower() == 'true'
|
||||
@@ -41,7 +41,7 @@ class Cleanup(SientiaMonitoring):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
storage_repository: StorageRepository,
|
||||
minio_repository: MinioRepository,
|
||||
logger: Logger,
|
||||
notification_handler: NotificationHandler,
|
||||
metrics_controller: MetricsController,
|
||||
@@ -50,13 +50,13 @@ class Cleanup(SientiaMonitoring):
|
||||
Initialize Cleanup activity.
|
||||
|
||||
Args:
|
||||
storage_repository: Repository for MinIO operations
|
||||
minio_repository: Repository for MinIO operations
|
||||
logger: Logger instance for observability
|
||||
notification_handler: Handler for sending notifications
|
||||
metrics_controller: Controller for metrics emission
|
||||
"""
|
||||
SientiaMonitoring.__init__(self, logger, notification_handler, metrics_controller)
|
||||
self.storage_repository = storage_repository
|
||||
self.minio_repository = minio_repository
|
||||
|
||||
# Configuration from environment variables
|
||||
self.retention_hours = RETENTION_HOURS
|
||||
@@ -113,15 +113,21 @@ class Cleanup(SientiaMonitoring):
|
||||
files_deleted = 0
|
||||
errors = []
|
||||
|
||||
# List objects in the specified bucket
|
||||
max_keys = self.max_keys_cleanup # Use environment variable for page size
|
||||
objects = self.storage_repository.list_bucket_objects(bucket_name, max_keys)
|
||||
# List objects in the specified bucket. MinioRepository applies BASE_PREFIX
|
||||
# internally; we request all objects under that prefix for this bucket.
|
||||
objects = await self.minio_repository.list_objects(
|
||||
prefix='',
|
||||
bucket=bucket_name,
|
||||
recursive=True,
|
||||
metadata=metadata,
|
||||
)
|
||||
|
||||
for obj_key in objects:
|
||||
files_scanned += 1
|
||||
|
||||
# Extract timestamp from filename
|
||||
match = self.minio_timestamp_pattern.match(obj_key)
|
||||
# Extract timestamp from the filename portion of the object key
|
||||
filename = obj_key.split('/')[-1]
|
||||
match = self.minio_timestamp_pattern.match(filename)
|
||||
if not match:
|
||||
self.debug(f'Skipping file without timestamp pattern: {obj_key}', metadata)
|
||||
continue
|
||||
@@ -137,10 +143,14 @@ class Cleanup(SientiaMonitoring):
|
||||
files_deleted += 1
|
||||
else:
|
||||
try:
|
||||
self.storage_repository.delete_file(bucket_name, obj_key)
|
||||
await self.minio_repository.delete_file(
|
||||
object_name=obj_key,
|
||||
bucket=None,
|
||||
metadata=metadata,
|
||||
)
|
||||
self.info(f'Deleted stale file: {obj_key}', metadata)
|
||||
files_deleted += 1
|
||||
except OSError as e:
|
||||
except Exception as e: # noqa: BLE001
|
||||
error_msg = f'Failed to delete {obj_key}: {str(e)}'
|
||||
errors.append(error_msg)
|
||||
self.error(error_msg, metadata)
|
||||
|
||||
@@ -12,18 +12,20 @@ with workflow.unsafe.imports_passed_through():
|
||||
import traceback
|
||||
from typing import Any
|
||||
|
||||
import pandas as pd
|
||||
from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler
|
||||
from sientia_do.notifications.models import NotificationLevel
|
||||
from sientia_do.observability.logger import Logger
|
||||
from sientia_do.observability.metrics_controller import MetricsController
|
||||
from sientia_do.observability.sientia_monitoring import SientiaMonitoring
|
||||
from sientia_do.repository.minio_repository import MinioRepository
|
||||
from sientia_model.model_repository.mlflow_repository import SientiaMLflowRepository
|
||||
from sientia_model.model_repository.plugin_store import PluginStore
|
||||
|
||||
from model_manager.metrics import ACTIVITY_EXECUTION_TOTAL, WORKFLOW_EXECUTION_TOTAL
|
||||
from model_manager.utils.exceptions import ModelTrainingError
|
||||
from model_manager.utils.models.train_model_params import TrainModelParams
|
||||
from model_manager.utils.repository.model_repository import ModelRepository
|
||||
from model_manager.utils.repository.storage_repository import StorageRepository
|
||||
from model_manager.utils.repository.training_repository import TrainingRepository
|
||||
from model_manager.utils.repository.data_manager_repository import DataManagerRepository
|
||||
|
||||
|
||||
class Training(SientiaMonitoring):
|
||||
@@ -38,8 +40,9 @@ class Training(SientiaMonitoring):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model_repository: ModelRepository,
|
||||
storage_repository: StorageRepository,
|
||||
mlflow_repository: SientiaMLflowRepository,
|
||||
plugin_store: PluginStore,
|
||||
minio_repository: MinioRepository,
|
||||
logger: Logger,
|
||||
notification_handler: NotificationHandler,
|
||||
metrics_controller: MetricsController,
|
||||
@@ -52,9 +55,10 @@ class Training(SientiaMonitoring):
|
||||
notification_handler: Handler for sending notifications
|
||||
"""
|
||||
SientiaMonitoring.__init__(self, logger, notification_handler, metrics_controller)
|
||||
self.training_repository = TrainingRepository(logger)
|
||||
self.model_repository = model_repository
|
||||
self.storage_repository = storage_repository
|
||||
self.data_manager_repository = DataManagerRepository(logger)
|
||||
self.mlflow_repository = mlflow_repository
|
||||
self.plugin_store = plugin_store
|
||||
self.minio_repository = minio_repository
|
||||
|
||||
@activity.defn(name='validate_train_params')
|
||||
async def validate_train_params(self, input_data: dict[str, Any]) -> TrainModelParams:
|
||||
@@ -120,8 +124,8 @@ class Training(SientiaMonitoring):
|
||||
|
||||
This activity orchestrates the ML training pipeline:
|
||||
1. Validate input parameters.
|
||||
2. Train the model via TrainingRepository.
|
||||
3. Perform post-training calculations.
|
||||
2. Prepare data via DataManagerRepository.
|
||||
3. Train the model and compute metrics.
|
||||
|
||||
Args:
|
||||
input_data: Training configuration containing:
|
||||
@@ -130,41 +134,98 @@ class Training(SientiaMonitoring):
|
||||
- train_params (TrainModelParams | dict): Training parameters.
|
||||
|
||||
Returns:
|
||||
dict: Keys `run_name` and `run_dir` when training and saving succeed.
|
||||
dict: Key `run_name` when training and saving succeed.
|
||||
|
||||
Raises:
|
||||
ValueError: If input validation fails.
|
||||
Exception: If training fails (after sending notification).
|
||||
"""
|
||||
metadata = input_data.get('metadata', {})
|
||||
metadata = input_data.get('metadata')
|
||||
train_params = input_data['train_params']
|
||||
|
||||
if isinstance(train_params, dict):
|
||||
train_params = TrainModelParams.from_dict(train_params)
|
||||
# type: ignore[assignment]
|
||||
|
||||
|
||||
model_trained = False
|
||||
model_saved = False
|
||||
metrics_status = 'success'
|
||||
|
||||
try:
|
||||
with self.storage_repository.fetch_file(
|
||||
train_params.bucket_name, train_params.file_name
|
||||
) as uploaded_file:
|
||||
train_result = self.training_repository.train(uploaded_file, train_params)
|
||||
# Download training file bytes from MinIO
|
||||
train_bytes = await self.minio_repository.download_file(
|
||||
object_name=train_params.file_name,
|
||||
bucket=train_params.bucket_name,
|
||||
metadata=metadata,
|
||||
)
|
||||
|
||||
train_result = self.training_repository.after_train_calculation(
|
||||
train_params, train_result
|
||||
# Download optional validation file bytes from the same bucket
|
||||
val_bytes: bytes | None = None
|
||||
validation_name = getattr(train_params, 'validation_file_name', None)
|
||||
if validation_name is not None:
|
||||
val_bytes = await self.minio_repository.download_file(
|
||||
object_name=validation_name,
|
||||
bucket=train_params.bucket_name,
|
||||
metadata=metadata,
|
||||
)
|
||||
|
||||
model_trained = True
|
||||
train_result = self.model_repository.save_model(train_result)
|
||||
model_saved = True
|
||||
train_result = self.data_manager_repository.prepare_training_data(
|
||||
train_file_bytes=train_bytes,
|
||||
validation_file_bytes=val_bytes,
|
||||
params=train_params,
|
||||
metadata=metadata,
|
||||
)
|
||||
|
||||
return {
|
||||
'run_name': train_result.run_name,
|
||||
'run_dir': train_result.run_dir,
|
||||
}
|
||||
wrapper = await self.plugin_store.get_model(
|
||||
model_name=train_params.model_name,
|
||||
force_download=False,
|
||||
opt_params={},
|
||||
model_kwargs={},
|
||||
data_model_kwargs={},
|
||||
metadata=metadata,
|
||||
)
|
||||
|
||||
train_df = pd.concat([train_result.x_train, train_result.y_train], axis=1)
|
||||
val_df = pd.concat([train_result.x_test, train_result.y_test], axis=1)
|
||||
|
||||
wrapper.train(
|
||||
train_data=train_df,
|
||||
val_data=val_df,
|
||||
target=train_params.target_variable,
|
||||
)
|
||||
|
||||
# Generate predictions using the trained wrapper
|
||||
transformed_train, _ = wrapper.transform(train_df)
|
||||
transformed_val, _ = wrapper.transform(val_df)
|
||||
|
||||
y_train_pred_df, _ = wrapper.predict({}, transformed_train)
|
||||
y_val_pred_df, _ = wrapper.predict({}, transformed_val)
|
||||
|
||||
# Use the first column of the prediction DataFrame as the target prediction
|
||||
train_result.y_train_pred = y_train_pred_df.iloc[:, 0]
|
||||
train_result.y_pred = y_val_pred_df.iloc[:, 0]
|
||||
|
||||
train_result = self.data_manager_repository.compute_regression_metrics(
|
||||
train_params,
|
||||
train_result,
|
||||
)
|
||||
|
||||
model_trained = True
|
||||
|
||||
async with self.mlflow_repository.start_run(
|
||||
model_name=train_params.model_name,
|
||||
run_name=None,
|
||||
experiment_name=train_params.experiment_name,
|
||||
tags=None,
|
||||
metadata=metadata,
|
||||
) as run_info:
|
||||
wrapper.store_model(name=train_params.model_name)
|
||||
|
||||
model_saved = True
|
||||
|
||||
return {
|
||||
'run_name': run_info.run_name or run_info.run_id,
|
||||
}
|
||||
except Exception as e: # noqa: BLE001
|
||||
metrics_status = 'error'
|
||||
|
||||
@@ -205,7 +266,6 @@ class Training(SientiaMonitoring):
|
||||
Args:
|
||||
input_data: Cleanup configuration containing:
|
||||
- metadata (dict): Workflow execution metadata.
|
||||
- run_dir (str): Temporary directory to remove.
|
||||
- bucket_name (str): MinIO bucket of the uploaded file.
|
||||
- file_name (str): MinIO object key to delete.
|
||||
|
||||
@@ -213,19 +273,21 @@ class Training(SientiaMonitoring):
|
||||
Exception: If cleanup fails (after sending notification).
|
||||
"""
|
||||
metadata = input_data.get('metadata', {})
|
||||
run_dir = input_data.get('run_dir', '')
|
||||
bucket_name = input_data.get('bucket_name', '')
|
||||
file_name = input_data.get('file_name', '')
|
||||
metrics_status = 'success'
|
||||
|
||||
try:
|
||||
self.model_repository.cleanup_run_directory(run_dir)
|
||||
self.storage_repository.delete_file(bucket_name, file_name)
|
||||
await self.minio_repository.delete_file(
|
||||
object_name=file_name,
|
||||
bucket=bucket_name,
|
||||
metadata=metadata,
|
||||
)
|
||||
except Exception as e: # noqa: BLE001
|
||||
metrics_status = 'error'
|
||||
|
||||
error_msg = (
|
||||
f'Error cleaning up resources - Run directory: {run_dir}, '
|
||||
'Error cleaning up resources - '
|
||||
f'File: {bucket_name}/{file_name}, Error: {str(e)}'
|
||||
)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user