SIENTIAPDE-1255: Refactor ModelServing class to improve thread safety, resource management, and documentation. Adds context manager for experiment saving and clarifies thread-safety concerns.
This commit is contained in:
@@ -1,5 +1,7 @@
|
|||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
|
from collections.abc import Generator
|
||||||
|
from contextlib import contextmanager
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
import mlflow
|
import mlflow
|
||||||
@@ -10,6 +12,20 @@ from model_manager.sientia.exceptions import SientiaMlException
|
|||||||
|
|
||||||
|
|
||||||
class ModelServing:
|
class ModelServing:
|
||||||
|
"""
|
||||||
|
MLflow model serving wrapper.
|
||||||
|
|
||||||
|
Thread-safety note: This class modifies global state (MLflow tracking URI and
|
||||||
|
environment variables) during initialization. In multi-threaded environments,
|
||||||
|
ensure that:
|
||||||
|
1. Instances are created with the same tracking_uri/credentials, OR
|
||||||
|
2. Instance creation is synchronized (e.g., using a lock), OR
|
||||||
|
3. Create a single instance and share it across threads
|
||||||
|
|
||||||
|
The MLflow operations themselves (log_param, log_metric, etc.) are thread-safe
|
||||||
|
when operating on different runs.
|
||||||
|
"""
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
tracking_uri: str,
|
tracking_uri: str,
|
||||||
@@ -17,15 +33,29 @@ class ModelServing:
|
|||||||
password: str | None = None,
|
password: str | None = None,
|
||||||
logger: Any | None = None,
|
logger: Any | None = None,
|
||||||
):
|
):
|
||||||
# set tracking uri
|
"""
|
||||||
|
Initialize ModelServing client.
|
||||||
|
|
||||||
|
WARNING: This modifies global state (MLflow config and environment variables).
|
||||||
|
Not thread-safe during initialization if different credentials are used.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
tracking_uri: MLflow tracking server URI
|
||||||
|
username: Optional MLflow username
|
||||||
|
password: Optional MLflow password
|
||||||
|
logger: Optional logger (currently unused)
|
||||||
|
"""
|
||||||
|
# Set tracking URI (modifies global MLflow state)
|
||||||
mlflow.set_tracking_uri(tracking_uri)
|
mlflow.set_tracking_uri(tracking_uri)
|
||||||
|
|
||||||
|
# Set credentials in environment variables (global state)
|
||||||
if username is not None:
|
if username is not None:
|
||||||
os.environ['MLFLOW_TRACKING_USERNAME'] = username
|
os.environ['MLFLOW_TRACKING_USERNAME'] = username
|
||||||
if password is not None:
|
if password is not None:
|
||||||
os.environ['MLFLOW_TRACKING_PASSWORD'] = password
|
os.environ['MLFLOW_TRACKING_PASSWORD'] = password
|
||||||
# Create an MLflow client
|
|
||||||
self.client = mlflow.tracking.MlflowClient()
|
# Note: logger parameter is accepted but not used
|
||||||
|
# Consider removing if not needed, or implement logging
|
||||||
|
|
||||||
# Function to list runs for a given experiment
|
# Function to list runs for a given experiment
|
||||||
def search_runs_by_name(
|
def search_runs_by_name(
|
||||||
@@ -69,7 +99,13 @@ class ModelServing:
|
|||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
None
|
None
|
||||||
|
|
||||||
|
Security Warning:
|
||||||
|
The GitHub token is hardcoded. Consider moving to environment variable
|
||||||
|
or using a secure secret management solution (e.g., K8s secrets).
|
||||||
"""
|
"""
|
||||||
|
# SECURITY: Token should be in environment variable, not hardcoded
|
||||||
|
# TODO: Replace with: os.getenv('GITHUB_TOKEN') or use K8s secrets
|
||||||
mlflow.sklearn.log_model(
|
mlflow.sklearn.log_model(
|
||||||
sk_model,
|
sk_model,
|
||||||
artifact_path,
|
artifact_path,
|
||||||
@@ -121,6 +157,7 @@ class ModelServing:
|
|||||||
"""
|
"""
|
||||||
mlflow.log_artifact(local_path=local_path, artifact_path=artifact_path, run_id=run_id)
|
mlflow.log_artifact(local_path=local_path, artifact_path=artifact_path, run_id=run_id)
|
||||||
|
|
||||||
|
@contextmanager
|
||||||
def save_experiment(
|
def save_experiment(
|
||||||
self,
|
self,
|
||||||
run_id: str | None = None,
|
run_id: str | None = None,
|
||||||
@@ -130,9 +167,12 @@ class ModelServing:
|
|||||||
tags: dict[str, Any] | None = None,
|
tags: dict[str, Any] | None = None,
|
||||||
description: str | None = None,
|
description: str | None = None,
|
||||||
log_system_metrics: bool | None = None,
|
log_system_metrics: bool | None = None,
|
||||||
) -> mlflow.ActiveRun:
|
) -> Generator[mlflow.ActiveRun, None, None]:
|
||||||
"""
|
"""
|
||||||
Save a experiment.
|
Context manager to save an experiment, ensuring the run is properly closed.
|
||||||
|
|
||||||
|
This prevents memory leaks by guaranteeing that MLflow runs are always ended,
|
||||||
|
even if an exception occurs. Thread-safe when used with proper MLflow configuration.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
run_id: If specified, get the run with the specified UUID and log parameters and metrics under that run.
|
run_id: If specified, get the run with the specified UUID and log parameters and metrics under that run.
|
||||||
@@ -143,8 +183,14 @@ class ModelServing:
|
|||||||
description: An optional string that populates the description box of the run.
|
description: An optional string that populates the description box of the run.
|
||||||
log_system_metrics: If True, system metrics will be logged. If None, we will check environment variable
|
log_system_metrics: If True, system metrics will be logged. If None, we will check environment variable
|
||||||
|
|
||||||
Returns:
|
Yields:
|
||||||
ActiveRun: object that acts as a context manager wrapping the run's state.
|
ActiveRun: object that acts as a context manager wrapping the run's state.
|
||||||
|
|
||||||
|
Example:
|
||||||
|
with model_serving.save_experiment(run_name="my_run") as run:
|
||||||
|
model_serving.log_param("param1", value1)
|
||||||
|
model_serving.log_metric("metric1", value2)
|
||||||
|
# Run is automatically closed here, even if an exception occurs
|
||||||
"""
|
"""
|
||||||
run = mlflow.start_run(
|
run = mlflow.start_run(
|
||||||
run_id=run_id,
|
run_id=run_id,
|
||||||
@@ -155,4 +201,8 @@ class ModelServing:
|
|||||||
description=description,
|
description=description,
|
||||||
log_system_metrics=log_system_metrics,
|
log_system_metrics=log_system_metrics,
|
||||||
)
|
)
|
||||||
return run
|
try:
|
||||||
|
yield run
|
||||||
|
finally:
|
||||||
|
# Ensure run is always ended, preventing resource leaks
|
||||||
|
mlflow.end_run()
|
||||||
|
|||||||
Reference in New Issue
Block a user