SIENTIAPDE-1646

Update dependencies and refactor MLFlow activities

- Replaced direct GitHub dependencies in `requirements.txt` with specific versioned packages for `sientia_do` and `sientia_model`.
- Refactored imports in `activities.py` to streamline the code structure.
- Enhanced the `MLFlow` class in `mlflow.py` by introducing a method to resolve model aliases, improving flexibility in model lookups.
- Simplified shutdown logic in `worker.py` for better readability.
- Added new tests for MLFlow activities and improved existing test coverage for data handling and model retraining processes.
This commit is contained in:
vitor-aignosi
2026-05-06 15:31:25 -03:00
parent 1ce8b9d3a7
commit 424be007ef
12 changed files with 625 additions and 29 deletions

View File

@@ -10,14 +10,13 @@ with workflow.unsafe.imports_passed_through():
from sientia_model.model_repository.mlflow_repository import SientiaMLflowRepository
from sientia_model.model_repository.plugin_store import PluginStore
from laborious.utils.connectors_config import build_mlflow_config
from laborious.activities.api import API
from laborious.activities.gates import Gates
from laborious.activities.mlflow import MLFlow
from laborious.activities.model_metrics import ModelMetrics
from laborious.activities.opc import OPC
from laborious.activities.storage import Storage
from laborious.utils.connectors_config import build_mlflow_config
class Activities(Storage, MLFlow, Gates, OPC, ModelMetrics, API):

View File

@@ -12,7 +12,6 @@ with workflow.unsafe.imports_passed_through():
import numpy as np
import pandas as pd
from pandas import DataFrame, to_datetime
from sklearn.model_selection import train_test_split
from sientia_do.notifications.handlers import CoreNotificationHandler as NotificationHandler
from sientia_do.notifications.models import NotificationLevel
from sientia_do.observability.logger import Logger
@@ -27,6 +26,7 @@ with workflow.unsafe.imports_passed_through():
from sientia_do.utils.formatters import create_sample_dict
from sientia_model.model_repository.mlflow_repository import SientiaMLflowRepository
from sientia_model.model_repository.plugin_store import PluginStore
from sklearn.model_selection import train_test_split
from laborious.utils.dataframe_debug import build_dataframe_debug_message
from laborious.utils.models.minio_dataframe_payload import MinioDataFramePayload
@@ -53,6 +53,7 @@ class MLFlow(MinioManager):
"""
_MAX_DEBUG_DATAFRAME_ROWS = 100
_DEFAULT_MODEL_ALIAS = 'production'
def __init__(
self,
@@ -191,6 +192,21 @@ class MLFlow(MinioManager):
latest = max(versions, key=lambda v: int(v.version))
return str(latest.version)
def _resolve_model_alias(self, model_config: dict[str, Any] | None = None) -> str:
"""
Resolve which MLflow alias should be used for model lookup/promotion.
Args:
- model_config: Optional model configuration that may include ``alias``.
Return:
str: Alias name trimmed and normalized; defaults to ``production``.
"""
if not model_config:
return self._DEFAULT_MODEL_ALIAS
alias = str(model_config.get('alias', self._DEFAULT_MODEL_ALIAS)).strip()
return alias or self._DEFAULT_MODEL_ALIAS
@activity.defn(name='request_transform')
async def request_transform(self, input_data: dict[str, Any]) -> MinioDataFramePayload:
"""
@@ -217,6 +233,7 @@ class MLFlow(MinioManager):
model_name = input_data['model_name']
model_config = input_data.get('model_config', {})
model_alias = self._resolve_model_alias(model_config)
self._debug_dataframe('Raw input data:', data, metadata)
@@ -238,7 +255,7 @@ class MLFlow(MinioManager):
try:
wrapper = self.mlflow_repository.get_cached_model(
model_name=model_name,
alias='production',
alias=model_alias,
retention_minutes=model_config.get('retention_minutes', 0),
metadata=metadata,
)
@@ -317,6 +334,7 @@ class MLFlow(MinioManager):
model_name = input_data['model_name']
model_config = input_data.get('model_config', {})
model_alias = self._resolve_model_alias(model_config)
self._debug_dataframe('Input data for prediction:', data, metadata)
@@ -332,7 +350,7 @@ class MLFlow(MinioManager):
try:
wrapper = self.mlflow_repository.get_cached_model(
model_name=model_name,
alias='production',
alias=model_alias,
retention_minutes=model_config.get('retention_minutes', 0),
metadata=metadata,
)
@@ -490,15 +508,16 @@ class MLFlow(MinioManager):
}
try:
model_alias = self._resolve_model_alias(model_config)
mv_src = self.mlflow_repository._client.get_model_version_by_alias(
name=model_name,
alias='production',
alias=model_alias,
)
source_run_id = mv_src.run_id
wrapper = self.mlflow_repository.get_cached_model(
model_name=model_name,
alias='production',
alias=model_alias,
retention_minutes=0,
metadata=metadata,
)
@@ -595,10 +614,11 @@ class MLFlow(MinioManager):
version = self._resolve_model_version_for_run(run_id)
promote_alias = self._resolve_model_alias(input_data.get('model_config'))
self.mlflow_repository.promote_to_alias(
model_name=model_name,
version=version,
alias='production',
alias=promote_alias,
metadata=metadata,
)
@@ -643,9 +663,10 @@ class MLFlow(MinioManager):
model_name = input_data['model_name']
try:
model_alias = self._resolve_model_alias(input_data.get('model_config'))
mv = self.mlflow_repository._client.get_model_version_by_alias(
name=model_name,
alias='production',
alias=model_alias,
)
run_id = mv.run_id

View File

@@ -264,10 +264,8 @@ async def main():
logger.custom_error(f'An unhandled exception occurred: {e}', metadata)
exit_code = 1
finally:
if notification_handler:
notification_handler.shutdown()
if activities:
await activities.shutdown()
notification_handler.shutdown()
await activities.shutdown()
metrics.APP_UP.labels(pod_id=POD_ID).set(0) # Mark app as DOWN
sys.exit(exit_code)