visa-fraud-detection / src /monitoring /drift_detection.py
Dheeraj-13's picture
Upload folder using huggingface_hub
dda22ae verified
Raw
History Blame Contribute Delete
13.1 kB
"""
Production model monitoring and drift detection.
Implements:
- Population Stability Index (PSI) for feature distribution drift
- Model performance tracking over time
- Data quality monitoring
- Automated alerting when thresholds are breached
"""
import json
import logging
from datetime import datetime
from pathlib import Path
from typing import Optional
import numpy as np
import pandas as pd
from sklearn.metrics import (
average_precision_score,
f1_score,
precision_score,
recall_score,
roc_auc_score,
)
logger = logging.getLogger(__name__)
class DistributionDriftDetector:
"""Detects feature distribution drift using Population Stability Index (PSI).
PSI measures how much a feature's distribution has shifted between
the training (reference) period and the current (production) period.
PSI interpretation:
- < 0.1: No significant drift
- 0.1 - 0.2: Moderate drift (investigate)
- > 0.2: Significant drift (retrain model)
"""
def __init__(self, n_bins: int = 10):
self.n_bins = n_bins
self.reference_distributions: dict = {}
def fit(self, reference_df: pd.DataFrame, feature_names: list[str]) -> None:
"""Learn reference (training) distributions for each feature.
Args:
reference_df: Training data used as the baseline.
feature_names: Feature columns to monitor.
"""
for feature in feature_names:
values = reference_df[feature].dropna().values
if len(values) == 0:
continue
# Compute bin edges from reference data
bin_edges = np.percentile(values, np.linspace(0, 100, self.n_bins + 1))
bin_edges = np.unique(bin_edges) # Handle duplicate edges
# Compute reference histogram
hist, _ = np.histogram(values, bins=bin_edges)
proportions = hist / hist.sum()
proportions = np.clip(proportions, 1e-6, None) # Avoid log(0)
self.reference_distributions[feature] = {
"bin_edges": bin_edges,
"proportions": proportions,
}
logger.info("Reference distributions fitted for %d features", len(self.reference_distributions))
def compute_psi(self, production_df: pd.DataFrame, feature_names: Optional[list[str]] = None) -> dict:
"""Compute PSI for each feature between reference and production data.
Args:
production_df: Current production data.
feature_names: Features to check (defaults to all fitted features).
Returns:
Dict mapping feature names to PSI values and drift status.
"""
if feature_names is None:
feature_names = list(self.reference_distributions.keys())
results = {}
for feature in feature_names:
if feature not in self.reference_distributions:
continue
ref = self.reference_distributions[feature]
values = production_df[feature].dropna().values
if len(values) == 0:
results[feature] = {"psi": float("inf"), "status": "NO_DATA"}
continue
# Compute production histogram using reference bin edges
hist, _ = np.histogram(values, bins=ref["bin_edges"])
prod_proportions = hist / max(1, hist.sum())
prod_proportions = np.clip(prod_proportions, 1e-6, None)
# PSI formula: sum( (prod - ref) * ln(prod / ref) )
psi = np.sum(
(prod_proportions - ref["proportions"])
* np.log(prod_proportions / ref["proportions"])
)
status = "OK" if psi < 0.1 else "WARNING" if psi < 0.2 else "CRITICAL"
results[feature] = {
"psi": round(float(psi), 6),
"status": status,
}
return results
class ModelPerformanceMonitor:
"""Tracks model performance metrics over time in production.
Compares current performance against baseline thresholds
and generates alerts when degradation is detected.
"""
def __init__(self, config: Optional[dict] = None):
self.config = config or {}
self.performance_history: list[dict] = []
self.baseline_metrics: Optional[dict] = None
self.alert_threshold = self.config.get("performance_alert_threshold", 0.05)
def set_baseline(self, y_true: np.ndarray, y_proba: np.ndarray) -> None:
"""Set baseline performance from validation/test evaluation.
Args:
y_true: True labels.
y_proba: Predicted probabilities.
"""
self.baseline_metrics = self._compute_metrics(y_true, y_proba)
logger.info("Baseline metrics set: %s", self.baseline_metrics)
def evaluate(
self,
y_true: np.ndarray,
y_proba: np.ndarray,
threshold: float = 0.5,
period_label: Optional[str] = None,
) -> dict:
"""Evaluate model performance on a batch of production data.
Args:
y_true: Ground truth labels (may be delayed in production).
y_proba: Model predicted probabilities.
threshold: Decision threshold.
period_label: Label for this evaluation period.
Returns:
Dict with metrics, comparison to baseline, and alerts.
"""
current_metrics = self._compute_metrics(y_true, y_proba, threshold)
current_metrics["period"] = period_label or datetime.now().isoformat()
current_metrics["n_samples"] = len(y_true)
current_metrics["fraud_rate"] = float(y_true.mean())
# Compare to baseline
alerts = []
if self.baseline_metrics:
for metric in ["f1", "auc_roc", "avg_precision"]:
baseline_val = self.baseline_metrics.get(metric, 0)
current_val = current_metrics.get(metric, 0)
drop = baseline_val - current_val
if drop > self.alert_threshold:
alerts.append({
"metric": metric,
"baseline": round(baseline_val, 4),
"current": round(current_val, 4),
"drop": round(drop, 4),
"severity": "CRITICAL" if drop > 2 * self.alert_threshold else "WARNING",
})
current_metrics["alerts"] = alerts
self.performance_history.append(current_metrics)
if alerts:
logger.warning("Performance alerts detected: %s", json.dumps(alerts, indent=2))
else:
logger.info("Performance within acceptable bounds: F1=%.4f, AUC=%.4f",
current_metrics["f1"], current_metrics["auc_roc"])
return current_metrics
def _compute_metrics(
self,
y_true: np.ndarray,
y_proba: np.ndarray,
threshold: float = 0.5,
) -> dict:
"""Compute standard metrics."""
y_pred = (y_proba >= threshold).astype(int)
return {
"precision": float(precision_score(y_true, y_pred, zero_division=0)),
"recall": float(recall_score(y_true, y_pred, zero_division=0)),
"f1": float(f1_score(y_true, y_pred, zero_division=0)),
"auc_roc": float(roc_auc_score(y_true, y_proba)) if y_true.sum() > 0 else 0.0,
"avg_precision": float(average_precision_score(y_true, y_proba)) if y_true.sum() > 0 else 0.0,
"fraud_caught": int(((y_pred == 1) & (y_true == 1)).sum()),
"false_positives": int(((y_pred == 1) & (y_true == 0)).sum()),
"total_fraud": int(y_true.sum()),
"threshold": threshold,
}
def get_trend(self, metric: str = "f1", last_n: int = 10) -> list[dict]:
"""Get the trend of a specific metric over recent evaluation periods.
Args:
metric: Metric name to track.
last_n: Number of recent periods.
Returns:
List of (period, value) dicts.
"""
recent = self.performance_history[-last_n:]
return [
{"period": entry["period"], "value": entry.get(metric, 0)}
for entry in recent
]
def export_report(self, output_path: str) -> None:
"""Export monitoring report to JSON.
Args:
output_path: Path to save the report.
"""
report = {
"generated_at": datetime.now().isoformat(),
"baseline_metrics": self.baseline_metrics,
"latest_metrics": self.performance_history[-1] if self.performance_history else None,
"total_evaluations": len(self.performance_history),
"active_alerts": [
entry for entry in self.performance_history[-1].get("alerts", [])
] if self.performance_history else [],
"history": self.performance_history,
}
Path(output_path).parent.mkdir(parents=True, exist_ok=True)
with open(output_path, "w") as f:
json.dump(report, f, indent=2, default=str)
logger.info("Monitoring report exported to %s", output_path)
class DataQualityMonitor:
"""Monitors incoming data quality in production.
Checks for:
- Missing values exceeding thresholds
- Schema violations
- Value range violations
- Cardinality changes in categorical features
"""
def __init__(self):
self.reference_stats: Optional[dict] = None
def fit(self, reference_df: pd.DataFrame) -> None:
"""Learn reference data quality statistics."""
self.reference_stats = {
"null_rates": reference_df.isnull().mean().to_dict(),
"numeric_ranges": {
col: {
"min": float(reference_df[col].min()),
"max": float(reference_df[col].max()),
"mean": float(reference_df[col].mean()),
}
for col in reference_df.select_dtypes(include=[np.number]).columns
},
"categorical_cardinalities": {
col: int(reference_df[col].nunique())
for col in reference_df.select_dtypes(include=["object"]).columns
},
"columns": list(reference_df.columns),
}
logger.info("Data quality reference stats computed")
def check(self, production_df: pd.DataFrame) -> dict:
"""Run data quality checks on production data.
Args:
production_df: Incoming production data batch.
Returns:
Dict with quality check results and any issues found.
"""
issues = []
# Check for missing columns
if self.reference_stats:
missing_cols = set(self.reference_stats["columns"]) - set(production_df.columns)
if missing_cols:
issues.append({
"type": "MISSING_COLUMNS",
"severity": "CRITICAL",
"details": f"Missing columns: {missing_cols}",
})
# Check null rates
current_null_rates = production_df.isnull().mean().to_dict()
if self.reference_stats:
for col, rate in current_null_rates.items():
ref_rate = self.reference_stats["null_rates"].get(col, 0)
if rate > ref_rate + 0.1: # 10% increase threshold
issues.append({
"type": "NULL_RATE_SPIKE",
"severity": "WARNING",
"column": col,
"reference_rate": round(ref_rate, 4),
"current_rate": round(rate, 4),
})
# Check numeric ranges
if self.reference_stats:
for col, ref_range in self.reference_stats.get("numeric_ranges", {}).items():
if col not in production_df.columns:
continue
prod_min = float(production_df[col].min())
prod_max = float(production_df[col].max())
if prod_min < ref_range["min"] * 0.5 or prod_max > ref_range["max"] * 2:
issues.append({
"type": "RANGE_VIOLATION",
"severity": "WARNING",
"column": col,
"reference_range": [ref_range["min"], ref_range["max"]],
"current_range": [prod_min, prod_max],
})
result = {
"timestamp": datetime.now().isoformat(),
"n_records": len(production_df),
"n_columns": len(production_df.columns),
"issues": issues,
"status": "CRITICAL" if any(i["severity"] == "CRITICAL" for i in issues)
else "WARNING" if issues
else "OK",
}
if issues:
logger.warning("Data quality issues found: %d", len(issues))
else:
logger.info("Data quality check passed — %d records OK", len(production_df))
return result