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:
@@ -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):
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user