Refactor ModelMetrics to utilize DriftAnalysis for drift detection - Replaced ModelAnalysis with DriftAnalysis in the ModelMetrics class to enhance drift detection capabilities. - Updated method signatures and documentation to reflect the changes in target_name and return values. - Adjusted data handling to ensure compatibility with the new analysis methods and improved clarity in the drift metrics dataframe preparation.
402 lines
17 KiB
Python
402 lines
17 KiB
Python
from temporalio import activity, workflow
|
|
|
|
with workflow.unsafe.imports_passed_through():
|
|
import time
|
|
import traceback
|
|
import warnings
|
|
from typing import Any
|
|
|
|
import numpy as np
|
|
import pandas as pd
|
|
from pandas import DataFrame, Index, Series, to_datetime
|
|
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.temporal.constants import DATETIME_FORMAT_WITH_TZ
|
|
from sientia_model.analytics.drift_analysis import DriftAnalysis
|
|
|
|
from laborious import metrics
|
|
from laborious.utils.dataframe_debug import build_dataframe_debug_message
|
|
|
|
warnings.filterwarnings('ignore', category=RuntimeWarning, message='Degrees of freedom <= 0')
|
|
warnings.filterwarnings(
|
|
'ignore', category=RuntimeWarning, message='invalid value encountered in scalar divide'
|
|
)
|
|
|
|
|
|
class ModelMetrics(SientiaMonitoring):
|
|
"""
|
|
Metrics and statistical analysis activities for the Laborious pipeline.
|
|
|
|
This class centralizes drift/statistical computations and model-quality
|
|
aggregates used by scheduled workflows. Besides producing tabular outputs
|
|
for persistence, it also emits operational metrics (count, lag, error)
|
|
through ``SientiaMonitoring`` so execution health is observable in runtime.
|
|
"""
|
|
|
|
_MAX_DEBUG_DATAFRAME_ROWS = 100
|
|
|
|
def __init__(
|
|
self,
|
|
logger: Logger,
|
|
notification_handler: NotificationHandler,
|
|
metrics_controller: MetricsController,
|
|
):
|
|
SientiaMonitoring.__init__(self, logger, notification_handler, metrics_controller)
|
|
|
|
def close(self) -> None:
|
|
"""
|
|
Shutdown monitoring resources associated with model metrics activities.
|
|
|
|
This is invoked during worker teardown to flush/close metric controller
|
|
internals and prevent dangling telemetry tasks.
|
|
"""
|
|
SientiaMonitoring.shutdown(self)
|
|
|
|
def __del__(self):
|
|
self.close()
|
|
|
|
def _debug_dataframe(self, message: str, data: Any, metadata: dict[str, Any]) -> None:
|
|
"""
|
|
Log dataframe content only when row count is below the configured threshold
|
|
|
|
Args:
|
|
- message (str): Base log message to identify the dataframe in logs
|
|
- data (Any): Dataframe-like payload to be logged
|
|
- metadata (dict[str, Any]): Workflow metadata for contextual logging
|
|
"""
|
|
self.debug(
|
|
build_dataframe_debug_message(
|
|
message=message,
|
|
data=data,
|
|
max_rows=self._MAX_DEBUG_DATAFRAME_ROWS,
|
|
),
|
|
metadata,
|
|
)
|
|
|
|
def get_drift_metrics(
|
|
self,
|
|
reference_data: DataFrame,
|
|
target_data: DataFrame,
|
|
target_name: str,
|
|
reference_columns: Index,
|
|
drift_metrics: list[str],
|
|
chunk_period: str,
|
|
metadata: dict[str, Any],
|
|
) -> DataFrame:
|
|
"""
|
|
Compute univariate and multivariate drift outputs and merge them into one dataframe.
|
|
|
|
The method orchestrates three analysis stages (univariate drift,
|
|
multivariate drift, and dataframe projection), emitting lag/count/error
|
|
metrics for each stage independently so failures are attributable.
|
|
|
|
Args:
|
|
- reference_data (DataFrame): Baseline dataset representing expected behavior.
|
|
- target_data (DataFrame): Current analysis dataset to compare against reference.
|
|
- target_name (str): Target column name used by ``DriftAnalysis`` config.
|
|
- reference_columns (Index): Feature columns evaluated for drift.
|
|
- drift_metrics (list[str]): Enabled univariate methods.
|
|
- chunk_period (str): Time bucket granularity used by analysis methods.
|
|
- metadata (dict[str, Any]): Workflow metadata for logs and notifications.
|
|
|
|
Return:
|
|
DataFrame: Consolidated drift dataframe from ``get_drift_metrics_dataframe`` using
|
|
``method`` / ``value`` (and optional ``threshold``, ``drift_type``), ready for
|
|
activity-level formatting before Postgres export.
|
|
"""
|
|
|
|
config = {
|
|
'target': target_name,
|
|
'prediction': 'prediction',
|
|
'timestamp': 'timestamp',
|
|
'features': reference_columns,
|
|
}
|
|
|
|
drift_analysis = DriftAnalysis(config=config)
|
|
|
|
self._debug_dataframe(
|
|
f'Reference data: Size {reference_data.shape}', reference_data, metadata
|
|
)
|
|
|
|
self._debug_dataframe(f'Target data: Size {target_data.shape}', target_data, metadata)
|
|
|
|
core_labels = self.get_core_labels(metadata, operation_type='detect_univariate_drift')
|
|
start_time = time.time()
|
|
try:
|
|
univariate_drift = drift_analysis.detect_univariate_drift(
|
|
reference_df=reference_data,
|
|
analysis_df=target_data,
|
|
features=reference_columns,
|
|
timestamp_col=config['timestamp'],
|
|
methods=drift_metrics,
|
|
chunk_period=chunk_period,
|
|
)
|
|
except Exception as e:
|
|
self.error(f'Error detecting univariate drift: {e}', metadata)
|
|
self.emit_metric_sync(metric_object=metrics.MODEL_ANALYZE_ERROR_COUNT, tags=core_labels)
|
|
raise e
|
|
self.observe_lag_sync(start_time, metrics.MODEL_ANALYZE_LAG, core_labels)
|
|
self.emit_metric_sync(metric_object=metrics.MODEL_ANALYZE_COUNT, tags=core_labels)
|
|
|
|
core_labels = self.get_core_labels(metadata, operation_type='detect_multivariate_drift')
|
|
start_time = time.time()
|
|
try:
|
|
multivariate_drift = drift_analysis.detect_multivariate_drift(
|
|
reference_df=reference_data,
|
|
analysis_df=target_data,
|
|
features=reference_columns,
|
|
timestamp_col=config['timestamp'],
|
|
chunk_period=chunk_period,
|
|
)
|
|
except Exception as e:
|
|
self.error(f'Error detecting multivariate drift: {e}', metadata)
|
|
self.emit_metric_sync(metric_object=metrics.MODEL_ANALYZE_ERROR_COUNT, tags=core_labels)
|
|
raise e
|
|
self.observe_lag_sync(start_time, metrics.MODEL_ANALYZE_LAG, core_labels)
|
|
self.emit_metric_sync(metric_object=metrics.MODEL_ANALYZE_COUNT, tags=core_labels)
|
|
|
|
start_time = time.time()
|
|
core_labels = self.get_core_labels(metadata, operation_type='get_drift_metrics_dataframe')
|
|
try:
|
|
drift_df = drift_analysis.get_drift_metrics_dataframe(
|
|
univariate_drift=univariate_drift,
|
|
multivariate_drift=multivariate_drift,
|
|
)
|
|
except Exception as e:
|
|
self.error(f'Error getting drift metrics: {e}', metadata)
|
|
self.emit_metric_sync(metric_object=metrics.MODEL_ANALYZE_ERROR_COUNT, tags=core_labels)
|
|
raise e
|
|
self.observe_lag_sync(start_time, metrics.MODEL_ANALYZE_LAG, core_labels)
|
|
self.emit_metric_sync(metric_object=metrics.MODEL_ANALYZE_COUNT, tags=core_labels)
|
|
|
|
self._debug_dataframe(f'Drift dataframe: Size {drift_df.shape}', drift_df, metadata)
|
|
|
|
return drift_df
|
|
|
|
@staticmethod
|
|
def _to_naive_utc(series: Series) -> Series:
|
|
"""
|
|
Parse ``series`` as datetime and return a TZ-naive UTC copy.
|
|
|
|
``sientia_model.analytics.drift_analysis.DriftAnalysis`` preserves the
|
|
timezone of the input dataframe in its outputs, while target rows
|
|
loaded from PostgreSQL come in with ``+00:00``. Forcing both sides of
|
|
a comparison to TZ-naive UTC keeps ``isin`` / ``floor`` operations
|
|
deterministic regardless of how the analyzer constructs its
|
|
timestamps.
|
|
|
|
Args:
|
|
- series (Series): Input series containing datetime-parseable values.
|
|
|
|
Return:
|
|
Series: Datetime64 series with ``tz=None`` representing UTC instants.
|
|
"""
|
|
parsed = to_datetime(series)
|
|
if getattr(parsed.dt, 'tz', None) is not None:
|
|
parsed = parsed.dt.tz_convert('UTC').dt.tz_localize(None)
|
|
return parsed
|
|
|
|
@activity.defn(name='calculate_drift')
|
|
def calculate_drift(self, input_data: dict[str, Any]) -> list[dict]:
|
|
"""
|
|
Calculate drift metrics for a model.
|
|
|
|
Args:
|
|
input_data (dict[str, Any]): Input data containing:
|
|
- metadata (dict): Workflow execution metadata
|
|
- model_name (str): Name of the MLFlow model to calculate drift for
|
|
- reference_data (pd.DataFrame): Reference data for the model
|
|
- target_data (pd.DataFrame): Target data for calculating drift
|
|
- target_name (str): Name of the target column
|
|
- drift_metrics (list[str]): List of drift metrics to calculate
|
|
"""
|
|
metadata = input_data['metadata']
|
|
model_name = input_data['model_name']
|
|
model_id = input_data['model_id']
|
|
reference_raw_data = input_data['reference_data']
|
|
target_data = DataFrame(input_data['target_data'])
|
|
target_name = input_data['target_name']
|
|
drift_metrics = input_data['drift_metrics']
|
|
chunk_period = input_data['chunk_period']
|
|
|
|
if chunk_period not in ['min', 's']:
|
|
self.error(f'Invalid chunk period: {chunk_period}', metadata)
|
|
raise ValueError(f'Invalid chunk period: {chunk_period}, must be "min" or "s"')
|
|
|
|
self.info(f'Calculating drift for model {model_name}', metadata)
|
|
|
|
target_data = target_data.pivot(index='timestamp', columns='variable', values='value')
|
|
target_data['timestamp'] = target_data.index
|
|
# Keep timestamps as datetime: DriftAnalysis._chunk_dataframe relies on
|
|
# ``pd.Grouper(freq=...)`` which rejects string timestamp columns.
|
|
target_data['timestamp'] = to_datetime(target_data['timestamp'])
|
|
target_data = target_data.reset_index(drop=True)
|
|
target_data.dropna(inplace=True)
|
|
|
|
if reference_raw_data is not None:
|
|
self.info('Using reference data', metadata)
|
|
reference_data = DataFrame(reference_raw_data)
|
|
if 'timestamp' in reference_data.columns:
|
|
reference_data['timestamp'] = to_datetime(reference_data['timestamp'])
|
|
accurate = True
|
|
else:
|
|
# Get 30% first rows of target_data
|
|
self.warning('Using 30% first rows of target data as reference data', metadata)
|
|
target_data.sort_values(by='timestamp', ascending=True, inplace=True)
|
|
reference_data = target_data.head(int(len(target_data) * 0.3))
|
|
accurate = False
|
|
|
|
self.send_notification(
|
|
metadata=metadata,
|
|
notification_id='MODEL_METRICS_REFERENCE_DATA_WARNING',
|
|
message='Using 30% first rows of target data as reference data',
|
|
block='model_metrics',
|
|
level=NotificationLevel.WARNING,
|
|
attachment_content=reference_data.to_csv(),
|
|
)
|
|
|
|
reference_columns = reference_data.drop(
|
|
columns=[target_name, 'timestamp', 'target', 'prediction'], errors='ignore'
|
|
).columns
|
|
|
|
try:
|
|
drift_df = self.get_drift_metrics(
|
|
reference_data=reference_data,
|
|
target_data=target_data,
|
|
target_name=target_name,
|
|
reference_columns=reference_columns,
|
|
drift_metrics=drift_metrics,
|
|
chunk_period=chunk_period,
|
|
metadata=metadata,
|
|
)
|
|
except Exception as e:
|
|
self.error(f'Error getting drift metrics: {e}', metadata)
|
|
self.send_notification(
|
|
metadata=metadata,
|
|
notification_id='MODEL_METRICS_GET_DRIFT_METRICS_ERROR',
|
|
message=f'Error getting drift metrics: {e}',
|
|
block='model_metrics',
|
|
level=NotificationLevel.ERROR,
|
|
attachment_content=traceback.format_exc(),
|
|
)
|
|
return []
|
|
|
|
if drift_df.empty:
|
|
self.warning('No drift metrics found', metadata)
|
|
return []
|
|
|
|
# Defense-in-depth: drop chunks whose floored timestamp does not appear
|
|
# in the analysis window. ``DriftAnalysis`` already chunks only over
|
|
# ``analysis_df`` so this only excludes rows injected by upstream
|
|
# callers that pre-merge reference data into the result.
|
|
target_floor = self._to_naive_utc(target_data['timestamp']).dt.floor(chunk_period)
|
|
drift_floor = self._to_naive_utc(drift_df['timestamp']).dt.floor(chunk_period)
|
|
drift_df = drift_df[drift_floor.isin(target_floor)]
|
|
|
|
if drift_df.empty:
|
|
self.warning(
|
|
'No drift metrics found after dropping rows where timestamp is not in target data',
|
|
metadata,
|
|
)
|
|
return []
|
|
|
|
# Analyzer emits diagnostic columns that are not stored in ``sientia_data.drift_metrics``.
|
|
drift_df = drift_df.drop(columns=['threshold', 'drift_type'], errors='ignore')
|
|
|
|
drift_df['model_id'] = str(model_id)
|
|
drift_df['accurate'] = accurate
|
|
|
|
# ``timestamp`` is overridden with the most recent target instant so
|
|
# every persisted row shares a single business timestamp (the run's
|
|
# logical "now"), matching what downstream consumers expect.
|
|
latest_target_timestamp = self._to_naive_utc(target_data['timestamp']).max()
|
|
drift_df['timestamp'] = (
|
|
pd.Timestamp(latest_target_timestamp)
|
|
.tz_localize('UTC')
|
|
.strftime(DATETIME_FORMAT_WITH_TZ)
|
|
)
|
|
|
|
# ``chunk_start_date`` / ``chunk_end_date`` may carry nanosecond
|
|
# precision (beyond ``timestamptz`` microseconds), so serialize as ISO
|
|
# text for the ``text`` Postgres columns.
|
|
for column in ('chunk_start_date', 'chunk_end_date'):
|
|
drift_df[column] = drift_df[column].apply(
|
|
lambda value: pd.Timestamp(value).isoformat() if pd.notna(value) else None
|
|
)
|
|
|
|
self._debug_dataframe(f'Drift dataframe: Size {drift_df.shape}', drift_df, metadata)
|
|
|
|
return drift_df.to_dict(orient='records')
|
|
|
|
@activity.defn(name='calculate_simple_metrics')
|
|
def calculate_simple_metrics(self, input_data: dict[str, Any]) -> list[dict]:
|
|
"""
|
|
Calculate simple metrics for a model. Metrics available are:
|
|
- rmse
|
|
- mse
|
|
- mae
|
|
- r2
|
|
- accuracy
|
|
- precision
|
|
- recall
|
|
- f1
|
|
Args:
|
|
input_data (dict[str, Any]): Input data containing:
|
|
- metadata (dict): Workflow execution metadata
|
|
- model_id (str): ID of the MLFlow model
|
|
- target_data (pd.DataFrame): Target data for calculating metrics, containing target and prediction columns
|
|
- metrics (list[str]): List of metrics to calculate
|
|
Returns:
|
|
dict[Hashable, Any]: Dictionary containing the calculated metrics
|
|
"""
|
|
|
|
metadata = input_data['metadata']
|
|
model_id = input_data['model_id']
|
|
target_data = DataFrame(input_data['target_data'])
|
|
metrics = input_data['metrics']
|
|
interval_minutes = input_data['interval_minutes']
|
|
|
|
data_size = target_data.shape[0]
|
|
|
|
output_data = []
|
|
|
|
diff = target_data['target'] - target_data['prediction']
|
|
diff_squared = diff**2
|
|
|
|
self.info(f'Calculating simple metrics for model {model_id}: {metrics}', metadata)
|
|
|
|
for metric in metrics:
|
|
if metric == 'rmse':
|
|
output_data.append({'metric': 'rmse', 'value': np.sqrt(np.mean(diff_squared))})
|
|
elif metric == 'mse':
|
|
output_data.append({'metric': 'mse', 'value': np.mean(diff_squared)})
|
|
elif metric == 'mae':
|
|
output_data.append({'metric': 'mae', 'value': np.mean(np.abs(diff))})
|
|
elif metric == 'r2':
|
|
y_true = target_data['target']
|
|
y_mean = np.mean(y_true)
|
|
|
|
ss_res = np.sum(diff_squared)
|
|
ss_tot = np.sum((y_true - y_mean) ** 2)
|
|
|
|
# Evita divisão por zero
|
|
if ss_tot == 0:
|
|
r2_score = 0.0
|
|
else:
|
|
r2_score = 1 - (ss_res / ss_tot)
|
|
|
|
output_data.append({'metric': 'r2', 'value': r2_score})
|
|
|
|
data = DataFrame(output_data)
|
|
data['model_id'] = model_id
|
|
data['timestamp'] = target_data['timestamp'].max()
|
|
data['data_size'] = data_size
|
|
data['interval_minutes'] = interval_minutes
|
|
|
|
self._debug_dataframe(f'Simple metrics dataframe: Size {data.shape}', data, metadata)
|
|
|
|
return data.to_dict(orient='records')
|