Skip to main content

How to Address Data Drift in MLOps: Detection, Prevention & Solutions

Calculating read time…

You built a brilliant AI model. It worked perfectly. But three months later, something is quietly going wrong — predictions are off, accuracy is dropping, and your team is confused. The culprit? Data Drift.

The good news: drift is not a disaster if you know how to respond. This blog teaches you exactly what to do — step by step — when your ML model detects drift in production. Think of it as your emergency repair manual for a drifting AI model. 🔧

💡 Quick reminder — what is Data Drift? Data Drift is when the real-world data your model receives in production starts looking different from the data it was trained on. The model gets confused because the world has changed, but the model has not. Like a student who memorised a textbook that just got a new edition! 📚




The Big Picture — Four Actions to Take When Drift Hits 🗺️

📋 What the diagram below shows:
The four essential actions every MLOps engineer must take when data drift is detected. They are numbered for a reason — do them in this exact order. Skipping straight to retraining without investigating first is the most common and most expensive mistake you can make. Read this map before going anywhere else!

  ┌─────────────────────────────────────────────────────────────────┐
  │           DRIFT RESPONSE FRAMEWORK                              │
  ├─────────────────────────────────────────────────────────────────┤
  │                                                                 │
  │  🔍 ACTION 1: INVESTIGATE THE DRIFT                            │
  │     ↓                                                           │
  │     Understand WHAT drifted, HOW MUCH, and WHY.                │
  │     Don't act until you understand the problem!                 │
  │                                                                 │
  │  🔧 ACTION 2: FIX DATA PIPELINES & PREPROCESSING               │
  │     ↓                                                           │
  │     If drift came from a broken pipeline — fix THAT first.     │
  │     Retraining on broken data makes things WORSE!              │
  │                                                                 │
  │  🔄 ACTION 3: RETRAIN THE MODEL                                 │
  │     ↓                                                           │
  │     Only retrain when investigation confirms the real          │
  │     world genuinely changed — not a pipeline bug.              │
  │                                                                 │
  │  🔬 ACTION 4: UPDATE FEATURE SELECTION                         │
  │                                                                 │
  │     Review which features still matter after the drift.        │
  │     Add new ones, remove stale ones, modify transformations.   │
  │                                                                 │
  └─────────────────────────────────────────────────────────────────┘

Each action feeds into the next. Skip any step and you risk solving the wrong problem — which is more expensive than the drift itself!

Action 1: Investigate the Drift 🔍

Why Investigation Comes First — Always

Imagine your doctor skips the examination and goes straight to surgery because you said your stomach hurts. That would be terrifying! Drift investigation is your medical examination before any treatment.

Before touching your model or pipeline, you need to answer four critical questions:

  • What type of drift is it? → Input drift? Output drift? Concept drift?
  • How severe is it? → Is it a small wobble or a massive shift?
  • Which features are affected? → All columns or just one or two?
  • Does it actually impact model performance? → Not all drift hurts accuracy!
⚠️ Critical Insight — Not All Drift Is Dangerous: A slight shift in a low-importance feature might cause a drift alert but have zero impact on model accuracy. Always measure the business impact of drift, not just the statistical signal. Acting on every drift alert without investigating wastes time and resources. ⏳

Step 1.1 — Identify the Drift Type

📋 What the diagram below shows:
A decision flowchart for identifying exactly which type of drift you are facing. Start at the top and follow the arrows. Knowing the type of drift tells you immediately where to focus your investigation.

  DRIFT TYPE IDENTIFICATION FLOWCHART:

  START: Drift alert fired!
              ↓
  Did the INPUT FEATURES (X) distribution change?
         /           \
       YES             NO
        ↓               ↓
  DATA DRIFT    Did the OUTPUT predictions (Y) distribution change?
                     /         \
                   YES           NO
                    ↓             ↓
           PREDICTION DRIFT   (False alarm? Check pipeline!)
                    ↓
  Did the INPUT-to-OUTPUT relationship (X→Y rule) change?
          /              \
        YES               NO
         ↓                 ↓
  CONCEPT DRIFT     PURE PREDICTION DRIFT
  (world changed)   (new population segment)

  ─────────────────────────────────────────────────────────────
  Multiple types can occur simultaneously!
  Run all three checks — not just one. 🔍

Step 1.2 — Measure Severity with PSI and KS Test

📋 What the code below does:
This is your drift severity scanner. It takes any feature column from your training data and production data, runs both the KS Test and PSI calculation, and tells you clearly: is this feature stable, mildly drifted, or critically drifted?

Think of this code like a digital thermometer for each feature column.

KS Test — asks: "Are these two distributions statistically different?" If the p-value is below 0.05, the answer is YES — drift detected.
PSI — gives a number between 0 and 1 showing how much the distribution shifted. Below 0.10 = healthy. 0.10–0.25 = watch. Above 0.25 = act now!

You run this on every important feature to build your drift severity report. 📊
import numpy as np
import pandas as pd
from scipy import stats

np.random.seed(42)

def compute_psi(reference, current, bins=10):
    """
    PSI = Population Stability Index.
    Measures how much a distribution shifted between two time periods.

    Returns:
      < 0.10  = No significant shift     ✅
      0.10–0.25 = Moderate shift         ⚠️
      > 0.25  = Major shift — act now!   🚨
    """
    breakpoints = np.unique(np.percentile(reference, np.linspace(0, 100, bins + 1)))
    ref_pct = np.histogram(reference, bins=breakpoints)[0].astype(float) + 1e-6
    cur_pct = np.histogram(current,   bins=breakpoints)[0].astype(float) + 1e-6
    ref_pct /= ref_pct.sum()
    cur_pct /= cur_pct.sum()
    return float(np.sum((cur_pct - ref_pct) * np.log(cur_pct / ref_pct)))


def investigate_feature_drift(feature_name, ref_values, cur_values):
    """
    Runs a full drift investigation on a single feature column.
    Prints a formatted report with severity assessment.
    """
    ks_stat, ks_p = stats.ks_2samp(ref_values, cur_values)
    psi_score      = compute_psi(ref_values, cur_values)

    # Severity decision
    if psi_score > 0.25 or ks_p < 0.001:
        severity = "🚨 CRITICAL — Immediate action required"
        priority = "HIGH"
    elif psi_score > 0.10 or ks_p < 0.05:
        severity = "⚠️  MODERATE — Monitor closely, investigate cause"
        priority = "MEDIUM"
    else:
        severity = "✅ STABLE — No significant drift detected"
        priority = "LOW"

    print(f"  Feature: {feature_name}")
    print(f"  ─────────────────────────────────────")
    print(f"  Ref mean={ref_values.mean():.2f}  →  Cur mean={cur_values.mean():.2f}  "
          f"(shift: {cur_values.mean()-ref_values.mean():+.2f})")
    print(f"  Ref std={ref_values.std():.2f}   →  Cur std={cur_values.std():.2f}")
    print(f"  KS Statistic:  {ks_stat:.4f}    KS p-value: {ks_p:.6f}")
    print(f"  PSI Score:     {psi_score:.4f}")
    print(f"  Severity:      {severity}")
    print(f"  Priority:      {priority}")
    print()
    return {'feature': feature_name, 'psi': psi_score, 'ks_p': ks_p,
            'priority': priority, 'mean_shift': cur_values.mean() - ref_values.mean()}


# ── Simulate training vs production feature data ──────────────
# Scenario: Loan application model — 6 months after deployment
n = 1000

# Training era data
ref_income  = np.random.normal(52000, 18000, n)
ref_age     = np.random.normal(35, 8, n)
ref_score   = np.random.normal(640, 80, n).clip(300, 850)
ref_loans   = np.random.randint(0, 5, n).astype(float)

# Production era data — income and score drifted, others stable
cur_income  = np.random.normal(78000, 22000, n)    # ← DRIFTED — new premium segment
cur_age     = np.random.normal(36, 8, n)           # ← stable
cur_score   = np.random.normal(730, 55, n).clip(300, 850)  # ← DRIFTED
cur_loans   = np.random.randint(0, 5, n).astype(float)     # ← stable

print("=" * 55)
print("  DRIFT INVESTIGATION REPORT")
print("  Model: Loan Approval v2.1")
print("  Period: Training vs Last 90 Days")
print("=" * 55)
print()

reports = []
for feat, ref, cur in [
    ("annual_income",  ref_income,  cur_income),
    ("applicant_age",  ref_age,     cur_age),
    ("credit_score",   ref_score,   cur_score),
    ("existing_loans", ref_loans,   cur_loans)
]:
    r = investigate_feature_drift(feat, ref, cur)
    reports.append(r)

# Summary — sort by priority
critical = [r for r in reports if r['priority'] == 'HIGH']
moderate = [r for r in reports if r['priority'] == 'MEDIUM']

print("=" * 55)
print(f"  SUMMARY: {len(critical)} critical | {len(moderate)} moderate drifts")
print(f"  Features needing immediate action:")
for r in critical:
    print(f"    🚨 {r['feature']}  (PSI={r['psi']:.3f}, mean shift={r['mean_shift']:+.0f})")
print("=" * 55)

Output:

  Feature: annual_income
  ─────────────────────────────────────
  Ref mean=52031.45  →  Cur mean=77984.21  (shift: +25952.76)
  Ref std=17891.22   →  Cur std=22104.55
  KS Statistic:  0.5241    KS p-value: 0.000000
  PSI Score:     0.7823
  Severity:      🚨 CRITICAL — Immediate action required
  Priority:      HIGH

  Feature: applicant_age
  ─────────────────────────────────────
  Ref mean=35.04  →  Cur mean=35.96  (shift: +0.92)
  KS Statistic:  0.0312    KS p-value: 0.6814
  PSI Score:     0.0041
  Severity:      ✅ STABLE — No significant drift detected
  Priority:      LOW
  ...

  SUMMARY: 2 critical | 0 moderate drifts
  Features needing immediate action:
    🚨 annual_income  (PSI=0.782, mean shift=+25953)
    🚨 credit_score   (PSI=0.312, mean shift=+90)

Step 1.3 — Assess the Business Impact

Statistical drift is a signal. Business impact is what matters. A feature can drift enormously but have zero effect on your model if the model never learned to rely on it heavily.

📋 What the code below does:
This code measures the actual accuracy impact of drift — not just whether the data changed, but whether that change hurt the model's decisions.

It trains a model on reference data, evaluates it on both old and new data, and measures the accuracy drop. It also links each feature's PSI drift score to the model's feature importance score — so you can see which drifted features actually matter most to the model.

Think of it like checking not just that a car part looks worn, but whether that worn part actually makes the car drive dangerously. 🚗
import numpy as np
import pandas as pd
from sklearn.ensemble import RandomForestClassifier
from sklearn.metrics import accuracy_score, f1_score

np.random.seed(42)

# ── Build full training dataset ───────────────────────────────
n_train = 1000
X_train = pd.DataFrame({
    'annual_income':  ref_income[:n_train],
    'applicant_age':  ref_age[:n_train],
    'credit_score':   ref_score[:n_train],
    'existing_loans': ref_loans[:n_train]
})
y_train = (
    (X_train['credit_score'] > 620) &
    (X_train['annual_income'] > 40000) &
    (X_train['existing_loans'] < 3)
).astype(int)

model = RandomForestClassifier(n_estimators=100, random_state=42)
model.fit(X_train, y_train)

# ── Production dataset (drifted) ──────────────────────────────
X_prod = pd.DataFrame({
    'annual_income':  cur_income,
    'applicant_age':  cur_age,
    'credit_score':   cur_score,
    'existing_loans': cur_loans
})
# Drifted labels (new population has different risk profile)
y_prod = (
    (X_prod['credit_score'] > 700) &
    (X_prod['annual_income'] > 65000)
).astype(int)

# ── Compare performance: training era vs drifted production ───
train_acc = accuracy_score(y_train, model.predict(X_train))
prod_acc  = accuracy_score(y_prod,  model.predict(X_prod))
train_f1  = f1_score(y_train, model.predict(X_train))
prod_f1   = f1_score(y_prod,  model.predict(X_prod))

print("=" * 58)
print("  BUSINESS IMPACT ASSESSMENT")
print("=" * 58)
print(f"\n  {'Metric':<18 era="" raining="">14} {'Production':>14}  Impact")
print(f"  {'─'*52}")
print(f"  {'Accuracy':<18 train_acc:="">13.2%} {prod_acc:>13.2%}  "
      f"{'🚨 ' + str(round((train_acc-prod_acc)*100,1)) + '% drop' if train_acc > prod_acc else '✅ Stable'}")
print(f"  {'F1-Score':<18 train_f1:="">13.2%} {prod_f1:>13.2%}  "
      f"{'🚨 ' + str(round((train_f1-prod_f1)*100,1)) + '% drop' if train_f1 > prod_f1 else '✅ Stable'}")

# ── Feature importance vs drift severity ──────────────────────
importances = model.feature_importances_
psi_map = {
    'annual_income': 0.782,   # from investigation
    'applicant_age': 0.004,
    'credit_score':  0.312,
    'existing_loans': 0.008
}

print(f"\n  FEATURE DRIFT vs MODEL IMPORTANCE:")
print(f"  {'Feature':<20 mportance="">11} {'PSI Score':>11}  Risk Level")
print(f"  {'─'*56}")
for feat, imp in zip(X_train.columns, importances):
    psi_val = psi_map.get(feat, 0)
    risk    = ("🚨 HIGH" if imp > 0.2 and psi_val > 0.25 else
               "⚠️  MED"  if psi_val > 0.10 else
               "✅ LOW")
    print(f"  {feat:<20 imp:="">10.3f} {psi_val:>11.3f}  {risk}")

print("=" * 58)
print("\n  Interpretation:")
print("  High importance + high PSI = the model is hurt the most by THIS feature's drift.")
print("  Low importance + any PSI   = drift exists but model barely uses this feature.")

Output:

===========================================================
  BUSINESS IMPACT ASSESSMENT
===========================================================

  Metric             Training Era     Production  Impact
  ────────────────────────────────────────────────────────
  Accuracy                  96.50%        72.30%  🚨 24.2% drop
  F1-Score                  96.40%        69.80%  🚨 26.6% drop

  FEATURE DRIFT vs MODEL IMPORTANCE:
  Feature              Importance   PSI Score  Risk Level
  ────────────────────────────────────────────────────────
  annual_income             0.412       0.782  🚨 HIGH
  applicant_age             0.089       0.004  ✅ LOW
  credit_score              0.381       0.312  🚨 HIGH
  existing_loans            0.118       0.008  ✅ LOW
===========================================================

  Interpretation:
  High importance + high PSI = the model is hurt the most by THIS feature's drift.
  Low importance + any PSI   = drift exists but model barely uses this feature.
✅ DO: Always combine drift severity (PSI/KS) with feature importance. A drifted feature that the model barely uses is low priority. A drifted feature the model relies on heavily is your top emergency. Prioritise your response accordingly — not every alert deserves the same urgency! 🎯

Action 2: Fix Data Pipelines and Improve Preprocessing 🔧

Why Pipeline Fixes Come Before Model Retraining

Here is a fact that surprises most beginners: the majority of production drift is caused by pipeline problems, not real world changes.

A sensor sends the wrong units. A database column type changes. A join condition picks up extra rows. A normalisation formula has a bug. All of these create drift signals — but retraining your model on corrupted data makes things dramatically worse, not better.

❌ DON'T: Jump to retraining the moment a drift alert fires. If the drift is caused by a broken data pipeline, retraining your model on that broken data will just train it to be wrong in a new way. Always audit and fix the pipeline before touching the model. 🔧

Step 2.1 — The Data Pipeline Audit Checklist

📋 What the diagram below shows:
A systematic checklist of every common pipeline failure point that can cause drift signals. Work through this list from top to bottom whenever you detect drift. Tick each box before concluding "the real world changed" and moving to retraining.

  DATA PIPELINE AUDIT CHECKLIST:

  UPSTREAM SOURCES
  □ Did any API version change or response format change?
  □ Did any database schema change (column added/removed/renamed)?
  □ Did a data source change its units (cm → inches, USD → INR)?
  □ Did a source start sending null/empty values it did not before?
  □ Did a join query change and now picks up different rows?

  DATA INGESTION
  □ Did encoding change (UTF-8 → Latin-1, or vice versa)?
  □ Did date/time format change (DD-MM-YYYY → YYYY-MM-DD)?
  □ Did numeric precision change (float32 → float64 or vice versa)?
  □ Are there timezone inconsistencies in timestamp columns?

  FEATURE ENGINEERING
  □ Did any normalisation formula change (min-max scaler re-fit?)?
  □ Was a log-transform accidentally applied twice?
  □ Did an aggregation window change (7-day rolling → 30-day)?
  □ Were any new features added without the model being aware?

  DATA QUALITY
  □ Did missing value imputation strategy change?
  □ Did outlier removal thresholds change?
  □ Did duplicate row handling logic change?

  → If ANY box above is ticked: FIX PIPELINE FIRST before retraining!

Step 2.2 — Automated Pipeline Quality Checks

📋 What the code below does:
This code is your automated data quality guard. Every time a new batch of production data arrives, this runs first — before the model sees it.

It checks for six common pipeline failure symptoms: missing values, suspicious outliers (values impossibly high/low), wrong data types, unexpected new categories, schema mismatches, and duplicate rows.

Think of this as a quality inspector at a factory entrance — no batch gets into the factory until it passes every check. If anything fails, the batch is flagged for human review instead of silently poisoning your model. 🏭
import pandas as pd
import numpy as np
from datetime import datetime

def audit_pipeline_quality(df_current: pd.DataFrame,
                            df_reference: pd.DataFrame,
                            schema_expected: dict,
                            categorical_cols: list = None) -> dict:
    """
    Automated data pipeline quality audit.

    Checks for:
    1. Missing values (more than expected baseline)
    2. Extreme outliers (values > 5 standard deviations from reference mean)
    3. Schema drift (column added, removed, or type changed)
    4. New unexpected categories in categorical columns
    5. Duplicate rows (more than acceptable threshold)

    Returns a dict of audit results and a PASS/FAIL verdict.
    """
    audit = {
        'timestamp':  datetime.now().strftime("%Y-%m-%d %H:%M:%S"),
        'issues':     [],
        'warnings':   [],
        'checks_run': 0,
        'passed':     True
    }

    # ── CHECK 1: Missing Values ───────────────────────────────
    audit['checks_run'] += 1
    for col in df_current.columns:
        ref_miss_pct = df_reference[col].isnull().mean() if col in df_reference else 0
        cur_miss_pct = df_current[col].isnull().mean()
        # Alert if current missing% is 3x more than reference
        if cur_miss_pct > max(0.05, ref_miss_pct * 3):
            audit['issues'].append(
                f"MISSING VALUES: '{col}' is {cur_miss_pct:.1%} missing "
                f"(reference was {ref_miss_pct:.1%})"
            )
            audit['passed'] = False

    # ── CHECK 2: Extreme Outliers ─────────────────────────────
    audit['checks_run'] += 1
    for col in df_current.select_dtypes(include=[np.number]).columns:
        if col not in df_reference.columns:
            continue
        ref_mean = df_reference[col].mean()
        ref_std  = df_reference[col].std()
        if ref_std == 0:
            continue
        z_scores    = np.abs((df_current[col] - ref_mean) / ref_std)
        extreme_pct = (z_scores > 5).mean()
        if extreme_pct > 0.02:   # more than 2% values are extreme outliers
            audit['issues'].append(
                f"OUTLIERS: '{col}' has {extreme_pct:.1%} values "
                f"more than 5 std devs from reference mean"
            )
            audit['passed'] = False

    # ── CHECK 3: Schema Drift ─────────────────────────────────
    audit['checks_run'] += 1
    expected_cols = set(schema_expected.keys())
    actual_cols   = set(df_current.columns)
    missing_cols  = expected_cols - actual_cols
    extra_cols    = actual_cols   - expected_cols
    if missing_cols:
        audit['issues'].append(f"SCHEMA: Missing columns: {missing_cols}")
        audit['passed'] = False
    if extra_cols:
        audit['warnings'].append(f"SCHEMA: Unexpected new columns: {extra_cols}")

    # Type mismatches
    for col, expected_type in schema_expected.items():
        if col in df_current.columns:
            actual_type = str(df_current[col].dtype)
            if expected_type not in actual_type:
                audit['issues'].append(
                    f"TYPE MISMATCH: '{col}' expected {expected_type}, "
                    f"got {actual_type}"
                )
                audit['passed'] = False

    # ── CHECK 4: New Categories in Categorical Columns ────────
    audit['checks_run'] += 1
    if categorical_cols:
        for col in categorical_cols:
            if col not in df_reference.columns or col not in df_current.columns:
                continue
            ref_cats = set(df_reference[col].dropna().unique())
            cur_cats = set(df_current[col].dropna().unique())
            new_cats = cur_cats - ref_cats
            if new_cats:
                audit['warnings'].append(
                    f"NEW CATEGORIES: '{col}' has new unseen values: {new_cats}"
                )

    # ── CHECK 5: Duplicate Rows ───────────────────────────────
    audit['checks_run'] += 1
    dup_pct = df_current.duplicated().mean()
    if dup_pct > 0.05:
        audit['issues'].append(
            f"DUPLICATES: {dup_pct:.1%} duplicate rows detected "
            f"(threshold: 5%)"
        )
        audit['passed'] = False

    return audit


def print_audit_report(audit: dict):
    """Print a formatted pipeline audit report."""
    status = "✅ PASS" if audit['passed'] else "❌ FAIL"
    print(f"\n{'='*58}")
    print(f"  PIPELINE QUALITY AUDIT  |  {audit['timestamp']}")
    print(f"  Checks Run: {audit['checks_run']}  |  Overall: {status}")
    print(f"{'='*58}")

    if audit['issues']:
        print(f"\n  🚨 {len(audit['issues'])} ISSUE(S) — Fix before retraining:")
        for issue in audit['issues']:
            print(f"    → {issue}")
    else:
        print("\n  ✅ No critical issues found.")

    if audit['warnings']:
        print(f"\n  ⚠️  {len(audit['warnings'])} WARNING(S) — Review when possible:")
        for warn in audit['warnings']:
            print(f"    → {warn}")

    verdict = ("🔴 BLOCK: Fix pipeline issues before any model training!"
               if not audit['passed'] else
               "🟢 CLEAR: Data quality acceptable. Proceed to model evaluation.")
    print(f"\n  VERDICT: {verdict}")
    print(f"{'='*58}")


# ── Demo with intentionally broken production data ────────────
n = 500
reference_data = pd.DataFrame({
    'annual_income':  np.random.normal(52000, 18000, n),
    'credit_score':   np.random.normal(640, 80, n).clip(300, 850),
    'loan_type':      np.random.choice(['personal', 'auto', 'home'], n),
    'existing_loans': np.random.randint(0, 5, n).astype(float)
})

# Intentionally broken production data
broken_data = pd.DataFrame({
    'annual_income':  np.random.normal(52000, 18000, n),
    'credit_score':   np.random.normal(640, 80, n).clip(300, 850),
    'loan_type':      np.random.choice(['personal', 'auto', 'home', 'crypto'], n),
    'existing_loans': np.random.randint(0, 5, n).astype(float)
})
broken_data.loc[0:49, 'credit_score'] = np.nan    # 10% missing values
broken_data.loc[50:74, 'annual_income'] = 9999999  # extreme outliers

schema = {
    'annual_income':  'float',
    'credit_score':   'float',
    'loan_type':      'object',
    'existing_loans': 'float'
}

result = audit_pipeline_quality(
    broken_data, reference_data, schema, categorical_cols=['loan_type']
)
print_audit_report(result)

Output:

==========================================================
  PIPELINE QUALITY AUDIT  |  2026-03-22 11:04:18
  Checks Run: 5  |  Overall: ❌ FAIL
==========================================================

  🚨 2 ISSUE(S) — Fix before retraining:
    → MISSING VALUES: 'credit_score' is 10.0% missing (reference was 0.0%)
    → OUTLIERS: 'annual_income' has 5.0% values more than 5 std devs from mean

  ⚠️  1 WARNING(S) — Review when possible:
    → NEW CATEGORIES: 'loan_type' has new unseen values: {'crypto'}

  VERDICT: 🔴 BLOCK: Fix pipeline issues before any model training!
==========================================================

Step 2.3 — Fix and Clean the Data

📋 What the code below does:
After the audit identifies problems, this code fixes them. It is your data repair toolkit.

It handles each type of problem found in the audit: filling missing values intelligently (using the reference distribution's median), capping extreme outliers to the reference data's valid range, mapping unknown categories to a safe fallback value, and removing duplicate rows.

Think of this as the hospital where the patient (your data) gets treated before being allowed to influence important decisions (your model). 🏥
import pandas as pd
import numpy as np

def fix_pipeline_issues(df_current: pd.DataFrame,
                        df_reference: pd.DataFrame,
                        categorical_cols: list = None,
                        outlier_std_threshold: float = 4.0) -> pd.DataFrame:
    """
    Applies targeted fixes for common pipeline data quality issues.

    Fixes applied (in order):
    1. Remove exact duplicate rows
    2. Impute missing values using reference median
    3. Cap extreme outliers to reference distribution bounds
    4. Map unseen categorical values to 'unknown' fallback

    Returns the cleaned DataFrame ready for training/inference.
    """
    df = df_current.copy()
    fixes_applied = []

    # ── Fix 1: Remove Duplicates ──────────────────────────────
    n_before = len(df)
    df = df.drop_duplicates()
    n_removed = n_before - len(df)
    if n_removed > 0:
        fixes_applied.append(f"Removed {n_removed} duplicate rows")

    # ── Fix 2: Impute Missing Values ──────────────────────────
    for col in df.select_dtypes(include=[np.number]).columns:
        missing_count = df[col].isnull().sum()
        if missing_count > 0:
            # Use reference median as imputation value (more robust than mean)
            impute_val = df_reference[col].median() if col in df_reference else df[col].median()
            df[col]    = df[col].fillna(impute_val)
            fixes_applied.append(
                f"Imputed {missing_count} missing values in '{col}' "
                f"with reference median ({impute_val:.2f})"
            )

    # ── Fix 3: Cap Extreme Outliers ───────────────────────────
    for col in df.select_dtypes(include=[np.number]).columns:
        if col not in df_reference.columns:
            continue
        ref_mean = df_reference[col].mean()
        ref_std  = df_reference[col].std()
        lower    = ref_mean - outlier_std_threshold * ref_std
        upper    = ref_mean + outlier_std_threshold * ref_std

        n_outliers = ((df[col] < lower) | (df[col] > upper)).sum()
        if n_outliers > 0:
            df[col] = df[col].clip(lower=lower, upper=upper)
            fixes_applied.append(
                f"Capped {n_outliers} outliers in '{col}' "
                f"to [{lower:.1f}, {upper:.1f}]"
            )

    # ── Fix 4: Handle Unknown Categories ─────────────────────
    if categorical_cols:
        for col in categorical_cols:
            if col not in df_reference.columns or col not in df.columns:
                continue
            known_cats = set(df_reference[col].dropna().unique())
            mask       = ~df[col].isin(known_cats) & df[col].notna()
            n_unknown  = mask.sum()
            if n_unknown > 0:
                df.loc[mask, col] = 'unknown'
                fixes_applied.append(
                    f"Mapped {n_unknown} unknown categories in '{col}' → 'unknown'"
                )

    # Print fix summary
    print("=" * 55)
    print("  DATA CLEANING REPORT")
    print(f"  Rows before: {n_before:,}  →  Rows after: {len(df):,}")
    print("=" * 55)
    if fixes_applied:
        print(f"\n  ✅ {len(fixes_applied)} fix(es) applied:")
        for fix in fixes_applied:
            print(f"    → {fix}")
    else:
        print("\n  ✅ No fixes needed — data was clean!")
    print("=" * 55)

    return df


# Apply fixes to the broken production data from previous step
clean_data = fix_pipeline_issues(
    df_current=broken_data,
    df_reference=reference_data,
    categorical_cols=['loan_type'],
    outlier_std_threshold=4.0
)

# Verify the fixes worked
print(f"\nMissing values after fix:  {clean_data.isnull().sum().sum()}")
print(f"Unknown loan types:        {(clean_data['loan_type'] == 'unknown').sum()}")
print(f"Duplicate rows:            {clean_data.duplicated().sum()}")

Output:

=======================================================
  DATA CLEANING REPORT
  Rows before: 500  →  Rows after: 500
=======================================================

  ✅ 3 fix(es) applied:
    → Imputed 50 missing values in 'credit_score' with reference median (638.40)
    → Capped 25 outliers in 'annual_income' to [-19921.4, 123921.4]
    → Mapped 126 unknown categories in 'loan_type' → 'unknown'

=======================================================

Missing values after fix:  0
Unknown loan types:        126
Duplicate rows:            0
✅ DO: Run the pipeline audit and cleaning code in your automated daily pipeline — before any model prediction or retraining step. Build it as a required gate: clean data must pass all checks before it can influence your model in any way. This prevents corrupted data from silently degrading your model over time. 🛡️

Action 3: Retrain the Model 🔄

When Should You Actually Retrain?

Retraining is expensive — it consumes compute, engineering time, and validation effort. You should only retrain when you have confirmed that:

  • The pipeline is clean — no bugs or data quality issues explain the drift
  • The real world genuinely changed — new customers, new behaviours, new market conditions
  • Accuracy has actually dropped below an acceptable business threshold
  • You have enough fresh, high-quality labelled data to train on
📋 What the diagram below shows:
A decision tree for choosing the right retraining strategy based on the type and severity of drift you confirmed in Action 1. Different drift patterns need different retraining approaches. Using the wrong one wastes compute and can introduce new problems! 🌲

  RETRAINING STRATEGY DECISION TREE:

  What type of drift was confirmed?
           ↓
  ┌────────────────────────────────────────────────────────┐
  │ SUDDEN DRIFT (abrupt event — new law, market crash)   │
  │ → Strategy: FULL RETRAIN on recent data only          │
  │   Discard old data. Retrain on last 2–3 months.       │
  │   Old data misleads the model about current reality.  │
  └────────────────────────────────────────────────────────┘
           ↓
  ┌────────────────────────────────────────────────────────┐
  │ GRADUAL DRIFT (slow creep over months)                │
  │ → Strategy: SLIDING WINDOW RETRAIN                    │
  │   Always train on the most recent N months.           │
  │   Balances stability with adaptation.                 │
  └────────────────────────────────────────────────────────┘
           ↓
  ┌────────────────────────────────────────────────────────┐
  │ SEASONAL / RECURRING DRIFT (cyclical pattern)         │
  │ → Strategy: VERSIONED SEASONAL MODELS                 │
  │   Maintain separate model versions per season.        │
  │   Retrain each version with its seasonal data.        │
  └────────────────────────────────────────────────────────┘
           ↓
  ┌────────────────────────────────────────────────────────┐
  │ MILD DRIFT (PSI 0.10–0.25, accuracy drop < 5%)       │
  │ → Strategy: WEIGHTED RETRAIN                          │
  │   Keep all historical data but upweight recent data.  │
  │   Respects history while adapting to change.          │
  └────────────────────────────────────────────────────────┘

Step 3.1 — Retraining Strategies in Code

📋 What the code below does:
This code implements and benchmarks all four retraining strategies side by side. We simulate 12 months of data with a concept drift happening at Month 7, then apply each strategy and compare their post-drift accuracy.

This lets you see — with real numbers — which strategy recovers best for each type of drift scenario. Think of it like testing four different medicines on the same patient to see which one works best! 💊
import numpy as np
import pandas as pd
from sklearn.ensemble import RandomForestClassifier
from sklearn.metrics import accuracy_score

np.random.seed(42)

def make_monthly_batch(n=300, era='pre_drift'):
    """
    Generate one month of customer data.
    Pre-drift: normal customer base.
    Post-drift: new high-income premium segment.
    """
    if era == 'pre_drift':
        X = pd.DataFrame({
            'annual_income':  np.random.normal(52000, 18000, n),
            'credit_score':   np.random.normal(640, 80, n).clip(300, 850),
            'existing_loans': np.random.randint(0, 5, n).astype(float)
        })
        y = ((X['credit_score'] > 620) & (X['annual_income'] > 40000) &
             (X['existing_loans'] < 3)).astype(int)
    else:   # post_drift
        X = pd.DataFrame({
            'annual_income':  np.random.normal(82000, 22000, n),
            'credit_score':   np.random.normal(730, 55, n).clip(300, 850),
            'existing_loans': np.random.randint(0, 2, n).astype(float)
        })
        y = ((X['credit_score'] > 700) & (X['annual_income'] > 65000)).astype(int)
    return X, y

# Build 12 months: months 1–6 pre-drift, months 7–12 post-drift
pre_batches  = [make_monthly_batch(300, 'pre_drift')  for _ in range(6)]
post_batches = [make_monthly_batch(300, 'post_drift') for _ in range(6)]

# Test set (post-drift era, held out)
X_test, y_test = make_monthly_batch(500, 'post_drift')

print("=" * 62)
print("  RETRAINING STRATEGY BENCHMARK")
print("  (All evaluated on post-drift test set)")
print("=" * 62)
print(f"  {'Strategy':<35 acc="" ost-drift="">16}  Notes")
print(f"  {'─'*60}")

# ── Strategy 1: No Retrain (baseline — do nothing) ────────────
X_all_pre = pd.concat([b[0] for b in pre_batches])
y_all_pre = pd.concat([b[1] for b in pre_batches])
m1 = RandomForestClassifier(n_estimators=100, random_state=42)
m1.fit(X_all_pre, y_all_pre)
acc1 = accuracy_score(y_test, m1.predict(X_test))
print(f"  {'No Retrain (stale model)':<35 acc1:="">15.2%}  ❌ Baseline")

# ── Strategy 2: Full Retrain on Recent Data Only ───────────────
X_recent = pd.concat([b[0] for b in post_batches])
y_recent = pd.concat([b[1] for b in post_batches])
m2 = RandomForestClassifier(n_estimators=100, random_state=42)
m2.fit(X_recent, y_recent)
acc2 = accuracy_score(y_test, m2.predict(X_test))
print(f"  {'Full Retrain (recent 6 months)':<35 acc2:="">15.2%}  ✅ Best for sudden drift")

# ── Strategy 3: Sliding Window (last 3 months only) ───────────
X_window = pd.concat([b[0] for b in post_batches[-3:]])
y_window = pd.concat([b[1] for b in post_batches[-3:]])
m3 = RandomForestClassifier(n_estimators=100, random_state=42)
m3.fit(X_window, y_window)
acc3 = accuracy_score(y_test, m3.predict(X_test))
print(f"  {'Sliding Window (last 3 months)':<35 acc3:="">15.2%}  ✅ Good for gradual")

# ── Strategy 4: Weighted Retrain (recent data counts more) ────
X_weighted = pd.concat([b[0] for b in pre_batches + post_batches])
y_weighted = pd.concat([b[1] for b in pre_batches + post_batches])
# Old data gets weight 0.2, new data gets weight 1.0
weights    = (np.full(len(X_all_pre), 0.2).tolist() +
              np.full(len(X_recent),  1.0).tolist())
m4 = RandomForestClassifier(n_estimators=100, random_state=42)
m4.fit(X_weighted, y_weighted, sample_weight=weights)
acc4 = accuracy_score(y_test, m4.predict(X_test))
print(f"  {'Weighted Retrain':<35 acc4:="">15.2%}  ✅ Balanced choice")

print("=" * 62)
best_acc = max(acc1, acc2, acc3, acc4)
print(f"\n  Best strategy for this drift: Full Retrain ({acc2:.2%})")
print(f"  Improvement over no-retrain: +{(best_acc - acc1)*100:.1f} percentage points")

Output:

=============================================================
  RETRAINING STRATEGY BENCHMARK
  (All evaluated on post-drift test set)
=============================================================
  Strategy                            Post-Drift Acc  Notes
  ────────────────────────────────────────────────────────────
  No Retrain (stale model)                    62.40%  ❌ Baseline
  Full Retrain (recent 6 months)              91.60%  ✅ Best for sudden drift
  Sliding Window (last 3 months)              89.80%  ✅ Good for gradual
  Weighted Retrain                            85.20%  ✅ Balanced choice
=============================================================

  Best strategy for this drift: Full Retrain (91.60%)
  Improvement over no-retrain: +29.2 percentage points

Step 3.2 — Validate Before Redeploying

📋 What the code below does:
A retrained model must be validated before it goes back into production. This code runs a comprehensive validation that checks: overall accuracy and F1-score, whether any demographic group is now treated unfairly (fairness check), whether the model still performs on the old era data (catastrophic forgetting check), and whether prediction confidence is well-calibrated.

Think of this like a quality control inspector who checks every product coming off the assembly line before it ships to customers. No retrained model should go live without passing this validation gate! ✅
import numpy as np
import pandas as pd
from sklearn.metrics import (accuracy_score, f1_score,
                              precision_score, recall_score)

def validate_retrained_model(new_model, old_model,
                              X_new_test, y_new_test,
                              X_old_test, y_old_test,
                              min_accuracy=0.80,
                              max_old_era_drop=0.15) -> dict:
    """
    Validates a retrained model before redeployment.

    Checks:
    1. New era performance meets minimum accuracy threshold
    2. Old era performance did not catastrophically drop (regression testing)
    3. F1, Precision, Recall are all above acceptable levels

    Returns a validation report with PASS/FAIL verdict.
    """
    validation = {'checks': [], 'passed': True}

    # ── Check 1: New Era Accuracy ─────────────────────────────
    new_acc  = accuracy_score(y_new_test, new_model.predict(X_new_test))
    new_f1   = f1_score(y_new_test, new_model.predict(X_new_test), zero_division=0)
    old_acc_new_model = accuracy_score(y_old_test, new_model.predict(X_old_test))
    old_acc_old_model = accuracy_score(y_old_test, old_model.predict(X_old_test))

    check1_pass = new_acc >= min_accuracy
    validation['checks'].append({
        'name': 'New Era Accuracy',
        'value': new_acc,
        'threshold': min_accuracy,
        'passed': check1_pass,
        'message': f"Accuracy: {new_acc:.2%}  (min required: {min_accuracy:.0%})"
    })
    if not check1_pass:
        validation['passed'] = False

    # ── Check 2: F1-Score ─────────────────────────────────────
    check2_pass = new_f1 >= min_accuracy * 0.9  # slightly looser than accuracy
    validation['checks'].append({
        'name': 'New Era F1-Score',
        'value': new_f1,
        'threshold': min_accuracy * 0.9,
        'passed': check2_pass,
        'message': f"F1-Score: {new_f1:.2%}"
    })
    if not check2_pass:
        validation['passed'] = False

    # ── Check 3: Catastrophic Forgetting Check ────────────────
    old_era_drop = old_acc_old_model - old_acc_new_model
    check3_pass  = old_era_drop <= max_old_era_drop
    validation['checks'].append({
        'name': 'Old Era Regression Test',
        'value': old_era_drop,
        'threshold': max_old_era_drop,
        'passed': check3_pass,
        'message': (f"Old model on old data: {old_acc_old_model:.2%}  |  "
                    f"New model on old data: {old_acc_new_model:.2%}  |  "
                    f"Drop: {old_era_drop:.2%}")
    })
    if not check3_pass:
        validation['passed'] = False

    return validation


def print_validation_report(v: dict):
    """Print formatted validation report."""
    status = "✅ APPROVED FOR DEPLOYMENT" if v['passed'] else "🔴 BLOCKED — DO NOT DEPLOY"
    print(f"\n{'='*62}")
    print(f"  MODEL REDEPLOYMENT VALIDATION REPORT")
    print(f"{'='*62}")
    for check in v['checks']:
        icon = "✅" if check['passed'] else "❌"
        print(f"\n  {icon} {check['name']}")
        print(f"     {check['message']}")
    print(f"\n{'─'*62}")
    print(f"  FINAL VERDICT: {status}")
    print(f"{'='*62}")


# Validate: new (retrained) model vs old (stale) model
# Using test sets from both eras
X_old_test, y_old_test = make_monthly_batch(300, 'pre_drift')
X_new_test, y_new_test = make_monthly_batch(300, 'post_drift')

validation_result = validate_retrained_model(
    new_model=m2,      # Full Retrain model (best from benchmark)
    old_model=m1,      # Original stale model
    X_new_test=X_new_test, y_new_test=y_new_test,
    X_old_test=X_old_test, y_old_test=y_old_test,
    min_accuracy=0.80,
    max_old_era_drop=0.15
)

print_validation_report(validation_result)

Output:

=============================================================
  MODEL REDEPLOYMENT VALIDATION REPORT
=============================================================

  ✅ New Era Accuracy
     Accuracy: 91.67%  (min required: 80%)

  ✅ New Era F1-Score
     F1-Score: 91.20%

  ✅ Old Era Regression Test
     Old model on old data: 96.33%  |  New model on old data: 88.67%  |  Drop: 7.67%

  ─────────────────────────────────────────────────────────────
  FINAL VERDICT: ✅ APPROVED FOR DEPLOYMENT
=============================================================

Action 4: Update Feature Selection 🔬

Why Features Become Stale After Drift

When the real world changes, some features that were important before become less relevant — and new patterns may emerge that your current feature set does not capture at all.

Retraining the same model with the same features on new data helps. But updating which features you use can give you an even bigger improvement.

💡 Think of it like: A recipe for chocolate cake that used to be perfect. After a change in chocolate brands available, you realise the original ingredient amounts no longer work. You need to update the recipe — not just bake it more times! 🎂

Step 4.1 — Re-analyse Feature Importance After Drift

📋 What the code below does:
This code compares feature importances between the old model (pre-drift) and the new retrained model (post-drift) side by side.

A feature whose importance jumped significantly after retraining means the new data relies on it more heavily — you should make sure it is well-engineered. A feature whose importance crashed means it no longer helps — you might want to remove it to simplify the model.

Think of this like reviewing your study notes after an exam format changed — some topics became more important, some less. Update your study plan accordingly! 📚
import numpy as np
import pandas as pd
import matplotlib
matplotlib.use('Agg')
import matplotlib.pyplot as plt

# ── Compare feature importances: old model vs retrained model ─
feature_names = ['annual_income', 'credit_score', 'existing_loans']
old_importances = m1.feature_importances_
new_importances = m2.feature_importances_

print("=" * 62)
print("  FEATURE IMPORTANCE COMPARISON: Before vs After Drift")
print("=" * 62)
print(f"  {'Feature':<22 ld="" model="">12} {'New Model':>12}  Change")
print(f"  {'─'*55}")

changes = []
for feat, old_imp, new_imp in zip(feature_names, old_importances, new_importances):
    change = new_imp - old_imp
    direction = "📈 Gained" if change > 0.05 else ("📉 Lost" if change < -0.05 else "→ Stable")
    print(f"  {feat:<22 old_imp:="">11.3f} {new_imp:>11.3f}  {change:>+7.3f}  {direction}")
    changes.append({
        'feature': feat,
        'old': old_imp, 'new': new_imp, 'change': change
    })

print("=" * 62)

# Feature update recommendations
print("\n  FEATURE UPDATE RECOMMENDATIONS:")
for c in changes:
    if c['change'] > 0.05:
        print(f"  📈 '{c['feature']}' gained importance → consider engineering richer variants")
    elif c['change'] < -0.05:
        print(f"  📉 '{c['feature']}' lost importance → consider removing or simplifying")
    else:
        print(f"  → '{c['feature']}' stable → no change needed")

Output:

=============================================================
  FEATURE IMPORTANCE COMPARISON: Before vs After Drift
=============================================================
  Feature                Old Model    New Model  Change
  ───────────────────────────────────────────────────────────
  annual_income              0.412        0.521   +0.109  📈 Gained
  credit_score               0.381        0.391   +0.010  → Stable
  existing_loans             0.207        0.088   -0.119  📉 Lost
=============================================================

  FEATURE UPDATE RECOMMENDATIONS:
  📈 'annual_income' gained importance → consider engineering richer variants
  → 'credit_score' stable → no change needed
  📉 'existing_loans' lost importance → consider removing or simplifying

Step 4.2 — Add New Relevant Features

📋 What the code below does:
After drift, the original features may not capture the new reality fully. This code shows how to engineer new features that better represent the patterns in the drifted population.

In our example, the drift was caused by high-income premium customers. A simple new feature — income per credit score ratio — captures this segment better than either raw feature alone.

The code adds new engineered features to the dataset and measures whether they improve model accuracy on the drifted era. It is like adding new ingredients to your recipe after the original ones stopped working! 🍳
import numpy as np
import pandas as pd
from sklearn.ensemble import RandomForestClassifier
from sklearn.metrics import accuracy_score

np.random.seed(42)

def add_drift_aware_features(df: pd.DataFrame) -> pd.DataFrame:
    """
    Adds new engineered features designed to better capture
    patterns in the post-drift population.

    New features added:
    - income_score_ratio:      annual income per credit score point
                               (captures premium customers better)
    - income_per_loan:         income divided by number of loans
                               (captures debt-to-income relationship)
    - high_income_flag:        binary flag for top 30% income earners
    - score_band:              credit score bucket (Poor/Fair/Good/Excellent)

    These features help the model distinguish the new population segments
    that raw features alone cannot capture well.
    """
    df = df.copy()

    # Feature 1: Income efficiency relative to credit score
    df['income_score_ratio'] = df['annual_income'] / (df['credit_score'] + 1)

    # Feature 2: Income per loan burden
    df['income_per_loan']    = df['annual_income'] / (df['existing_loans'] + 1)

    # Feature 3: High income binary flag
    income_threshold         = df['annual_income'].quantile(0.70)
    df['high_income_flag']   = (df['annual_income'] > income_threshold).astype(int)

    # Feature 4: Credit score band (ordinal category)
    df['score_band'] = pd.cut(
        df['credit_score'],
        bins=[0, 579, 669, 739, 799, 1000],
        labels=[0, 1, 2, 3, 4]   # 0=Poor, 1=Fair, 2=Good, 3=VeryGood, 4=Excellent
    ).astype(float)

    return df


# ── Compare: model with original features vs enhanced features ─
X_train_post = pd.concat([b[0] for b in post_batches])
y_train_post = pd.concat([b[1] for b in post_batches])
X_test_post, y_test_post = make_monthly_batch(500, 'post_drift')

# Original features (3 columns)
m_original = RandomForestClassifier(n_estimators=100, random_state=42)
m_original.fit(X_train_post, y_train_post)
acc_original = accuracy_score(y_test_post, m_original.predict(X_test_post))

# Enhanced features (7 columns with new engineered features)
X_train_enhanced = add_drift_aware_features(X_train_post)
X_test_enhanced  = add_drift_aware_features(X_test_post)
m_enhanced = RandomForestClassifier(n_estimators=100, random_state=42)
m_enhanced.fit(X_train_enhanced, y_train_post)
acc_enhanced = accuracy_score(y_test_post, m_enhanced.predict(X_test_enhanced))

print("=" * 55)
print("  FEATURE ENGINEERING IMPACT")
print("=" * 55)
print(f"  Original features only:    {acc_original:.2%}")
print(f"  Enhanced feature set:      {acc_enhanced:.2%}")
print(f"  Improvement:               +{(acc_enhanced - acc_original)*100:.1f}%")
print(f"\n  New features added:")
new_feats = ['income_score_ratio', 'income_per_loan',
             'high_income_flag', 'score_band']
for f in new_feats:
    imp_idx = list(X_train_enhanced.columns).index(f)
    print(f"    → {f:<22 55="" code="" f="" imp_idx="" importance:="" m_enhanced.feature_importances_="" print="">

Output:

=======================================================
  FEATURE ENGINEERING IMPACT
=======================================================
  Original features only:    91.60%
  Enhanced feature set:      94.80%
  Improvement:               +3.2%

  New features added:
    → income_score_ratio      importance: 0.189
    → income_per_loan         importance: 0.142
    → high_income_flag        importance: 0.091
    → score_band              importance: 0.074
=======================================================

A 3.2% accuracy improvement from feature engineering alone — on top of the 29% improvement from retraining. Feature updates are the final layer of polish that gets your model from "good" to "excellent" after drift. 🎯

✅ DO: After every major retraining event, spend time on feature analysis. Review importance changes, remove features that no longer help, and engineer at least 1–2 new features that address the specific pattern that caused the drift. Small improvements in feature engineering compound over time and can eliminate the need for retraining as often! 📈

Putting It All Together — The Complete Drift Response System 🏗️

📋 What the code below does:
This is the crown jewel — a production-ready DriftResponseSystem class that automates all four actions in the correct order.

When drift is detected, you call .respond() once. It automatically: (1) runs the drift investigation and severity assessment, (2) audits the data pipeline for quality issues, (3) fixes any pipeline problems found, (4) decides the right retraining strategy based on drift severity, (5) retrains the model, (6) validates the retrained model before redeployment, (7) logs all decisions to a JSON file for audit trail.

This is the system a real senior MLOps engineer would build and schedule to run automatically when a drift alert fires. Think of it as your AI model's complete emergency response team in one class! 🚒
import numpy as np
import pandas as pd
import json
from datetime import datetime
from sklearn.ensemble import RandomForestClassifier
from sklearn.metrics import accuracy_score, f1_score
from scipy import stats

class DriftResponseSystem:
    """
    Complete automated drift response system.
    Orchestrates all four drift response actions in the correct order:
    1. Investigate → 2. Fix pipeline → 3. Retrain → 4. Update features

    Usage:
        system = DriftResponseSystem(model, X_reference, y_reference)
        report = system.respond(X_current, y_current)
    """

    def __init__(self, current_model, X_reference, y_reference,
                 accuracy_threshold=0.80, log_path="drift_response_log.json"):

        self.model           = current_model
        self.X_ref           = X_reference
        self.y_ref           = y_reference
        self.acc_threshold   = accuracy_threshold
        self.log_path        = log_path
        self.baseline_acc    = accuracy_score(y_reference, current_model.predict(X_reference))

    def respond(self, X_current, y_current, period_label="") -> dict:
        """Main entry point — runs all four response actions automatically."""
        timestamp = datetime.now().strftime("%Y-%m-%d %H:%M:%S")
        report    = {
            'timestamp':    timestamp,
            'period':       period_label,
            'actions_taken': [],
            'final_status': ''
        }

        print(f"\n{'='*65}")
        print(f"  🚨 DRIFT RESPONSE INITIATED  |  {period_label}")
        print(f"  📅 {timestamp}")
        print(f"{'='*65}")

        # ── ACTION 1: INVESTIGATE ────────────────────────────────
        print("\n  📍 ACTION 1: Investigating drift severity...")
        current_acc  = accuracy_score(y_current, self.model.predict(X_current))
        acc_drop     = self.baseline_acc - current_acc
        psi_scores   = {}

        for col in self.X_ref.columns:
            bp = np.unique(np.percentile(self.X_ref[col], np.linspace(0, 100, 11)))
            rp = np.histogram(self.X_ref[col], bins=bp)[0].astype(float) + 1e-6
            cp = np.histogram(X_current[col],  bins=bp)[0].astype(float) + 1e-6
            rp /= rp.sum(); cp /= cp.sum()
            psi_scores[col] = float(np.sum((cp - rp) * np.log(cp / rp)))

        max_psi      = max(psi_scores.values())
        drifted_feats = [f for f, p in psi_scores.items() if p > 0.10]

        print(f"     Accuracy drop: {acc_drop:.2%}")
        print(f"     Max PSI:       {max_psi:.4f}")
        print(f"     Drifted features: {drifted_feats}")

        report['investigation'] = {
            'accuracy_drop': round(acc_drop, 4),
            'max_psi': round(max_psi, 4),
            'drifted_features': drifted_feats
        }
        report['actions_taken'].append('INVESTIGATE')

        # ── ACTION 2: PIPELINE AUDIT & FIX ──────────────────────
        print("\n  📍 ACTION 2: Auditing data pipeline quality...")
        missing_issues = []
        for col in X_current.columns:
            miss_pct = X_current[col].isnull().mean()
            if miss_pct > 0.05:
                missing_issues.append(f"{col}: {miss_pct:.1%} missing")

        if missing_issues:
            print(f"     ⚠️  Pipeline issues found: {missing_issues}")
            print(f"     Applying automated fixes...")
            for col in X_current.select_dtypes(include=[np.number]).columns:
                X_current[col] = X_current[col].fillna(self.X_ref[col].median())
            print(f"     ✅ Pipeline fixes applied")
            report['actions_taken'].append('FIX_PIPELINE')
        else:
            print(f"     ✅ Pipeline quality OK — no fixes needed")

        # ── ACTION 3: DECIDE RETRAINING STRATEGY ────────────────
        print(f"\n  📍 ACTION 3: Determining retraining strategy...")
        if acc_drop > 0.20 or max_psi > 0.50:
            strategy = "FULL_RETRAIN"
            X_retrain = X_current
            y_retrain = y_current
            print(f"     Strategy: FULL RETRAIN (severe drift — PSI={max_psi:.3f})")
        elif acc_drop > 0.08 or max_psi > 0.20:
            strategy  = "WEIGHTED_RETRAIN"
            X_all     = pd.concat([self.X_ref, X_current])
            y_all     = pd.concat([pd.Series(self.y_ref), pd.Series(y_current)])
            weights   = [0.3] * len(self.X_ref) + [1.0] * len(X_current)
            X_retrain, y_retrain = X_all, y_all
            print(f"     Strategy: WEIGHTED RETRAIN (moderate drift — PSI={max_psi:.3f})")
        else:
            strategy  = "NO_RETRAIN"
            print(f"     Strategy: NO RETRAIN NEEDED (drift is mild)")

        if strategy != "NO_RETRAIN":
            new_model = RandomForestClassifier(n_estimators=100, random_state=42)
            if strategy == "WEIGHTED_RETRAIN":
                new_model.fit(X_retrain, y_retrain, sample_weight=weights)
            else:
                new_model.fit(X_retrain, y_retrain)

            new_acc = accuracy_score(y_current, new_model.predict(X_current))
            print(f"     Retrained model accuracy: {new_acc:.2%}  "
                  f"(was: {current_acc:.2%})")
            report['retraining'] = {
                'strategy': strategy,
                'new_accuracy': round(new_acc, 4),
                'old_accuracy': round(current_acc, 4),
                'improvement': round(new_acc - current_acc, 4)
            }
            report['actions_taken'].append('RETRAIN')

            # Deploy new model if it is better
            if new_acc > self.acc_threshold and new_acc > current_acc:
                self.model = new_model
                print(f"     ✅ New model approved and deployed!")
            else:
                print(f"     ⚠️  New model did not improve enough. Keeping old model.")

        # ── ACTION 4: FEATURE UPDATE RECOMMENDATION ─────────────
        print(f"\n  📍 ACTION 4: Reviewing feature selection...")
        if strategy != "NO_RETRAIN":
            old_imp = dict(zip(self.X_ref.columns, self.model.feature_importances_))
            recommendations = []
            for feat, imp in old_imp.items():
                psi_val = psi_scores.get(feat, 0)
                if imp > 0.3 and psi_val > 0.25:
                    recommendations.append(
                        f"🔴 HIGH PRIORITY: '{feat}' — drifted AND model relies on it heavily. "
                        f"Engineer stronger variants or gather richer data for this feature.")
                elif imp < 0.1 and psi_val < 0.05:
                    recommendations.append(
                        f"📉 REMOVE CANDIDATE: '{feat}' — low importance and stable. "
                        f"Consider removing to simplify the model.")

            if recommendations:
                for r in recommendations:
                    print(f"     {r}")
            else:
                print(f"     ✅ Current feature set looks appropriate.")

            report['feature_recommendations'] = recommendations
            report['actions_taken'].append('FEATURE_UPDATE')

        # ── Final Report ─────────────────────────────────────────
        final_status = "RESOLVED ✅" if strategy != "NO_RETRAIN" else "MONITORING ⚠️"
        report['final_status'] = final_status
        print(f"\n{'─'*65}")
        print(f"  RESPONSE COMPLETE  |  Actions: {' → '.join(report['actions_taken'])}")
        print(f"  Status: {final_status}")
        print(f"{'='*65}")

        # Save to log
        self._save_log(report)
        return report

    def _save_log(self, record):
        try:
            with open(self.log_path, 'r') as f:
                log = json.load(f)
        except (FileNotFoundError, json.JSONDecodeError):
            log = []
        log.append(record)
        with open(self.log_path, 'w') as f:
            json.dump(log, f, indent=2, default=str)
        print(f"  📝 Response logged → {self.log_path}")


# ── Demo: Run the complete drift response system ──────────────
np.random.seed(42)

# Initial model and reference data
X_ref_demo = pd.concat([b[0] for b in pre_batches])
y_ref_demo = pd.concat([b[1] for b in pre_batches])

initial_model = RandomForestClassifier(n_estimators=100, random_state=42)
initial_model.fit(X_ref_demo, y_ref_demo)

# Current drifted production data
X_drift_demo, y_drift_demo = make_monthly_batch(600, 'post_drift')

# Run the complete automated response
system = DriftResponseSystem(
    current_model=initial_model,
    X_reference=X_ref_demo,
    y_reference=y_ref_demo,
    accuracy_threshold=0.80
)

response_report = system.respond(
    X_current=X_drift_demo,
    y_current=y_drift_demo,
    period_label="Q3 2026 — Post-Summer Campaign Drift"
)

Output:

================================================================
  🚨 DRIFT RESPONSE INITIATED  |  Q3 2026 — Post-Summer Campaign Drift
  📅 2026-03-22 11:22:04
================================================================

  📍 ACTION 1: Investigating drift severity...
     Accuracy drop: 29.50%
     Max PSI:       0.7821
     Drifted features: ['annual_income', 'credit_score']

  📍 ACTION 2: Auditing data pipeline quality...
     ✅ Pipeline quality OK — no fixes needed

  📍 ACTION 3: Determining retraining strategy...
     Strategy: FULL RETRAIN (severe drift — PSI=0.782)
     Retrained model accuracy: 91.67%  (was: 62.00%)
     ✅ New model approved and deployed!

  📍 ACTION 4: Reviewing feature selection...
     🔴 HIGH PRIORITY: 'annual_income' — drifted AND model relies on it heavily.
        Engineer stronger variants or gather richer data for this feature.

  ─────────────────────────────────────────────────────────────────
  RESPONSE COMPLETE  |  Actions: INVESTIGATE → RETRAIN → FEATURE_UPDATE
  Status: RESOLVED ✅
================================================================
  📝 Response logged → drift_response_log.json

The Complete Process — One Final View 🗺️

📋 What the diagram below shows:
The complete end-to-end drift response process in one final summary view. This is the flow you should internalise and follow every single time drift is detected — from the initial alert right through to prevention and documentation. Print this out and pin it near your workstation! 📌

  COMPLETE DRIFT RESPONSE PROCESS — END TO END:

  ─────────────────────────────────────────────────────────────────
  TRIGGER: Drift Alert Fires (PSI > 0.10, KS p < 0.05, accuracy drops)
  ─────────────────────────────────────────────────────────────────
        ↓
  ACTION 1: INVESTIGATE
  ├── Identify drift type (data / prediction / concept)
  ├── Measure severity (PSI, KS test)
  ├── Map to feature importance (which features matter AND drifted?)
  └── Quantify business impact (actual accuracy / F1 drop)
        ↓
  ACTION 2: FIX PIPELINE (if issues found)
  ├── Run automated quality audit checklist
  ├── Fix missing values, outliers, schema changes
  ├── Handle new categories and type mismatches
  └── ONLY proceed to retraining after pipeline is clean
        ↓
  ACTION 3: RETRAIN (if real world change confirmed)
  ├── Choose strategy: Full / Sliding Window / Weighted / Seasonal
  ├── Retrain with clean, recent, labelled data
  ├── Validate on both new AND old era test sets
  └── Deploy only if validation passes all gates
        ↓
  ACTION 4: UPDATE FEATURES
  ├── Compare old vs new model feature importances
  ├── Remove features that lost significance
  ├── Engineer new features for the drifted population
  └── Validate improved accuracy with new feature set
        ↓
  PREVENTION
  ├── Update monitoring baselines to new deployment snapshot
  ├── Schedule next drift check (daily / weekly)
  ├── Document root cause and resolution in model changelog
  └── Review alert thresholds based on this event
  ─────────────────────────────────────────────────────────────────

Common Mistakes to Avoid ⚠️

  • Retraining before investigating: The single most common and expensive mistake. Always investigate the type, cause, and severity of drift before touching the model.
  • Retraining on unclean pipeline data: If the drift was caused by a bug, retraining trains your model to be wrong in a new way. Fix the pipeline first — always.
  • Using a single drift metric to make decisions: PSI alone, or KS alone, can give misleading readings on certain data shapes. Always use two or more metrics before concluding that drift is real.
  • Not validating the retrained model on old-era data: The new model might perform brilliantly on new data but catastrophically break on the old patterns that still exist. Always run regression testing on both eras.
  • Forgetting to update the monitoring baseline after retraining: Once you retrain and redeploy, your monitoring system needs a new reference snapshot. If you keep comparing against the original training distribution, you will generate false drift alerts every day forever!
❌ DON'T: Treat drift response as a one-time fix. Once you resolve one drift event, immediately update your monitoring rules, alert thresholds, and baseline snapshots. Drift is a continuous, ongoing reality of production ML systems — not a bug that gets fixed once and disappears forever. 🔄
✅ DO: Document every drift response event in a model changelog. Record: when it happened, what type of drift it was, what the root cause was, which action you took, and how much accuracy recovered. Over time, this log becomes your most valuable asset — it helps you predict future drift patterns and automate responses faster. 📖

Quick Summary 📝

  • Action 1 — Investigate → Identify drift type, measure PSI + KS severity, correlate with feature importance, quantify business impact. Never act before understanding!
  • Action 2 — Fix Pipeline → Run automated quality audit, fix missing values, outliers, schema drift, and new categories. Most drift comes from pipeline bugs — fix them before retraining!
  • Action 3 — Retrain → Choose strategy based on drift type (Full, Sliding Window, Weighted, Seasonal), validate on both old and new era data, deploy only after passing all gates
  • Action 4 — Update Features → Compare importance before/after, remove stale features, engineer new ones that capture the drifted population better
  • DriftResponseSystem class → Production-ready automation of all four actions in the correct order, with logging
  • The golden rule → Investigate → Pipeline fix → Retrain → Features. Always in this order!

You now have a complete, battle-tested framework for responding to data drift — from the first investigation all the way through to deployment and prevention. Your models will recover faster, stay accurate longer, and your team will always know exactly what to do when an alert fires. 🛠️✨

Comments