Monitoring & Drift Detection
Logging, métricas, data drift, concept drift, alertas y reentrenamiento.
Un modelo desplegado en producción no es un artefacto estático: el mundo que intenta predecir cambia constantemente. Los clientes cambian su comportamiento, los datos de entrada evolucionan, y la distribución de las características que entrenaron el modelo puede divergir de la distribución real que llega a la API. Sin monitoreo activo, un modelo puede degradarse silenciosamente durante semanas antes de que alguien note que las predicciones ya no son confiables.
Data drift (también llamado covariate shift) ocurre cuando la distribución de las features de entrada cambia respecto a la distribución de entrenamiento. Ejemplo: entrenaste un modelo de churn con clientes de 25-45 años; si la base de clientes envejece y ahora domina el rango 50-70, las distribuciones divergen. El modelo sigue 'funcionando' técnicamente pero produce predicciones poco calibradas para el nuevo perfil de cliente.
Concept drift es más sutil y más peligroso: es cuando la relación entre las features y el target cambia. El modelo aprendió que tenure alto implica baja probabilidad de churn. Si una crisis económica cambia ese patrón (clientes con tenure alto también churnan), el modelo falla aunque las features sean idénticas a las de entrenamiento. Detectar concept drift requiere labels recientes —lo cual es costoso— o métricas de negocio como proxy (tasa de churn real vs predicha).
Las métricas operacionales (latencia, throughput, error rate) son distintas de las métricas de modelo (accuracy, precision, recall, PSI, KL-divergence). Ambas deben monitorearse. La latencia p99 te dice si hay un problema de infraestructura. El Population Stability Index (PSI) te dice si la distribución de una feature ha cambiado significativamente. Un PSI > 0.2 es la señal de alarma estándar en la industria.
El logging estructurado (JSON) es la base de todo sistema de monitoreo: cada request debe loggearse con timestamp, features de entrada, predicción, latencia, y un request_id que permita trazar el flujo completo. Herramientas como CloudWatch, Datadog o Grafana leen estos logs estructurados y permiten construir dashboards y alertas sin parsear texto libre.
El pipeline de reentrenamiento debe ser automático y dispararse por condiciones medibles: PSI > umbral, accuracy < umbral, volumen de datos nuevo suficiente. Un reentrenamiento manual ad-hoc no es MLOps —es gestión de crisis. El objetivo es tener un sistema que detecte drift, dispare el reentrenamiento, evalúe el nuevo modelo contra el anterior (champion/challenger), y promueva el mejor automáticamente.
# Structured logging — every request
import logging, json, time
from uuid import uuid4
logger = logging.getLogger('churn_api')
logger.setLevel(logging.INFO)
def log_prediction(features: dict, score: float, latency_ms: float) -> None:
logger.info(json.dumps({
'event': 'prediction',
'request_id': str(uuid4()),
'features': features,
'churn_probability': score,
'latency_ms': round(latency_ms, 2),
}))
# Middleware to measure latency
from fastapi import Request
import time
async def latency_middleware(request: Request, call_next):
t0 = time.perf_counter()
response = await call_next(request)
latency = (time.perf_counter() - t0) * 1000
response.headers['X-Latency-Ms'] = str(round(latency, 2))
return response
# Data drift — Population Stability Index
import numpy as np
def compute_psi(expected: np.ndarray, actual: np.ndarray, bins: int = 10) -> float:
breakpoints = np.percentile(expected, np.linspace(0, 100, bins + 1))
breakpoints[0] = -np.inf
breakpoints[-1] = np.inf
e_pct = np.histogram(expected, breakpoints)[0] / len(expected)
a_pct = np.histogram(actual, breakpoints)[0] / len(actual)
e_pct = np.where(e_pct == 0, 1e-4, e_pct)
a_pct = np.where(a_pct == 0, 1e-4, a_pct)
return float(np.sum((a_pct - e_pct) * np.log(a_pct / e_pct)))
# Alert rule
PSI_THRESHOLD = 0.2
def check_drift(train_feature: np.ndarray, prod_feature: np.ndarray) -> None:
psi = compute_psi(train_feature, prod_feature)
if psi > PSI_THRESHOLD:
logger.warning(json.dumps({'event': 'drift_alert', 'psi': psi}))
trigger_retraining(reason=f'PSI={psi:.3f}')Debugging lab
Detecta y corrige el error en el código.
- 5.4.5.1
def check_drift(train_data: list[float], prod_data: list[float]) -> bool: train_mean = sum(train_data) / len(train_data) prod_mean = sum(prod_data) / len(prod_data) if abs(train_mean - prod_mean) > 10: return True return False
- 5.4.5.2
# Monitoreo de latencia import time latencies = [] @app.post('/predict') def predict(req: PredictRequest): t0 = time.time() score = use_case.execute(req) latency = time.time() - t0 latencies.append(latency) avg = sum(latencies) / len(latencies) if avg > 0.5: alert('High latency') return score
- 5.4.5.3
# Logging de predicciones @app.post('/predict') def predict(req: PredictRequest): score = use_case.execute(req) print(f'Prediction: {score.probability} for age={req.age}') return score
- 5.4.5.4
# Reentrenamiento automático def check_and_retrain(psi: float) -> None: if psi > 0.1: retrain_model() deploy_model('new_model.pkl')
- 5.4.5.5
# Ventana de monitoreo def compute_weekly_psi(feature_name: str) -> float: all_prod_data = db.query(f'SELECT {feature_name} FROM predictions') train_data = load_training_feature(feature_name) return compute_psi(train_data, all_prod_data)