Deployment and Production Monitoring
Deployment patterns, production monitoring, drift detection, and operational controls for public health AI. The material is maintained separately so each operational question has a stable, focused reference.
- Identify the evidence and controls relevant to this decision area
- Distinguish technical performance from operational and population impact
- Apply the included framework without extending claims beyond the cited evidence
Introduction
This focused reference is part of the broader Deployment and Production Monitoring overview. It preserves the detailed methods, examples, and exercises while reducing page size and improving direct navigation.
Deployment Strategies
Model Serving Options
Option 1: REST API (Most Common)
Use case: Real-time predictions, synchronous requests, language-agnostic integration
Baylor et al., 2017, KDD - “TFX: A TensorFlow-Based Production-Scale Machine Learning Platform” describes production ML serving.
Implementation with FastAPI:
from fastapi import FastAPI, HTTPException, Depends, Header
from pydantic import BaseModel, Field, validator
import mlflow.pyfunc
import pandas as pd
import numpy as np
from typing import List, Dict, Optional
import logging
from datetime import datetime
import time
import hashlib
# Configure logging
logging.basicConfig(
level=logging.INFO,
format='%(asctime)s - %(name)s - %(levelname)s - %(message)s'
)
logger = logging.getLogger(__name__)
# Initialize FastAPI app
app = FastAPI(
title="Sepsis Prediction API",
description="Real-time sepsis risk prediction for ICU patients using ML",
version="1.2.0",
docs_url="/api/docs",
redoc_url="/api/redoc"
)
# Global model variable (loaded at startup)
model = None
model_metadata = {}
# Load model at startup
@app.on_event("startup")
async def load_model():
global model, model_metadata
try:
model = mlflow.pyfunc.load_model("models:/sepsis_predictor/Production")
# Load metadata
from mlflow.tracking import MlflowClient
client = MlflowClient()
prod_versions = client.get_latest_versions("sepsis_predictor", stages=["Production"])
if prod_versions:
version = prod_versions[0]
model_metadata = {
'version': version.version,
'run_id': version.run_id,
'created_at': datetime.fromtimestamp(version.creation_timestamp / 1000).isoformat(),
'tags': version.tags
}
logger.info(f"[OK] Model loaded successfully - Version {model_metadata.get('version')}")
except Exception as e:
logger.error(f"[ERROR] Failed to load model: {e}")
raise
# Input schema with validation
class PatientData(BaseModel):
patient_id: str = Field(..., description="Unique patient identifier", min_length=1, max_length=50)
heart_rate: float = Field(..., ge=0, le=300, description="Heart rate (bpm)")
respiratory_rate: float = Field(..., ge=0, le=60, description="Respiratory rate (breaths/min)")
temperature: float = Field(..., ge=35.0, le=42.0, description="Body temperature (°C)")
systolic_bp: float = Field(..., ge=40, le=300, description="Systolic blood pressure (mmHg)")
white_blood_cell: float = Field(..., ge=0, le=100, description="WBC count (K/μL)")
lactate: float = Field(..., ge=0, le=20, description="Serum lactate (mmol/L)")
age: int = Field(..., ge=18, le=120, description="Patient age (years)")
@validator('patient_id')
def validate_patient_id(cls, v):
# Remove any non-alphanumeric characters
import re
if not re.match(r'^[A-Za-z0-9-_]+$', v):
raise ValueError('patient_id must contain only alphanumeric characters, hyphens, and underscores')
return v
class Config:
schema_extra = {
"example": {
"patient_id": "ICU-2024-001",
"heart_rate": 110.5,
"respiratory_rate": 24.0,
"temperature": 38.5,
"systolic_bp": 95.0,
"white_blood_cell": 15.2,
"lactate": 2.8,
"age": 67
}
}
# Output schema
class PredictionResponse(BaseModel):
patient_id: str
sepsis_risk: float = Field(..., ge=0, le=1, description="Probability of sepsis (0-1)")
risk_category: str = Field(..., description="Risk level: Low, Medium, or High")
confidence: float = Field(..., ge=0, le=1, description="Model confidence (0-1)")
model_version: str
timestamp: str
prediction_id: str = Field(..., description="Unique prediction identifier for audit trail")
class Config:
schema_extra = {
"example": {
"patient_id": "ICU-2024-001",
"sepsis_risk": 0.73,
"risk_category": "High",
"confidence": 0.89,
"model_version": "1.2.0",
"timestamp": "2024-03-15T14:30:00Z",
"prediction_id": "pred_abc123def456"
}
}
# API key authentication (simplified - use proper auth in production)
async def verify_api_key(x_api_key: str = Header(...)):
# In production: validate against database, implement rate limiting
valid_keys = {"demo_key_123"} # Replace with secure key management
if x_api_key not in valid_keys:
raise HTTPException(status_code=403, detail="Invalid API key")
return x_api_key
@app.get("/")
def root():
"""Root endpoint with API information"""
return {
"service": "Sepsis Prediction API",
"version": "1.2.0",
"status": "operational",
"endpoints": {
"predict": "/api/predict",
"batch": "/api/predict/batch",
"health": "/health",
"metrics": "/metrics"
},
"documentation": "/api/docs"
}
@app.get("/health")
def health_check():
"""
Health check endpoint for load balancer
Returns 200 if service is healthy, 503 otherwise
"""
try:
if model is None:
raise HTTPException(status_code=503, detail="Model not loaded")
# Quick inference test
test_input = pd.DataFrame([{
'heart_rate': 80,
'respiratory_rate': 16,
'temperature': 37.0,
'systolic_bp': 120,
'white_blood_cell': 8.0,
'lactate': 1.0,
'age': 50
}])
_ = model.predict(test_input)
return {
"status": "healthy",
"model_loaded": True,
"model_version": model_metadata.get('version'),
"timestamp": datetime.utcnow().isoformat()
}
except Exception as e:
logger.error(f"Health check failed: {e}")
raise HTTPException(status_code=503, detail=f"Service unhealthy: {str(e)}")
@app.post("/api/predict", response_model=PredictionResponse)
async def predict(
data: PatientData,
api_key: str = Depends(verify_api_key)
):
"""
Make real-time sepsis risk prediction for a single patient
**Risk Categories:**
- Low: Risk < 0.3
- Medium: 0.3 ≤ Risk < 0.7
- High: Risk ≥ 0.7
**Authentication:** Requires valid API key in X-API-Key header
"""
start_time = time.time()
try:
# Convert to DataFrame
input_df = pd.DataFrame([{
'heart_rate': data.heart_rate,
'respiratory_rate': data.respiratory_rate,
'temperature': data.temperature,
'systolic_bp': data.systolic_bp,
'white_blood_cell': data.white_blood_cell,
'lactate': data.lactate,
'age': data.age
}])
# Make prediction
prediction = model.predict(input_df)[0]
sepsis_risk = float(prediction)
# Categorize risk
if sepsis_risk < 0.3:
risk_category = "Low"
elif sepsis_risk < 0.7:
risk_category = "Medium"
else:
risk_category = "High"
# Estimate confidence (simplified - in production use proper uncertainty quantification)
# For ensemble models, use prediction variance
confidence = min(1.0, max(abs(sepsis_risk - 0.5) * 2, 0.5))
# Generate unique prediction ID for audit trail
prediction_id = hashlib.sha256(
f"{data.patient_id}{datetime.utcnow().isoformat()}".encode()
).hexdigest()[:16]
# Calculate latency
latency_ms = (time.time() - start_time) * 1000
# Log prediction
logger.info(
f"Prediction: patient={data.patient_id}, "
f"risk={sepsis_risk:.3f}, category={risk_category}, "
f"latency={latency_ms:.1f}ms"
)
# In production: Log to database for monitoring and audit
# log_prediction_to_db(prediction_id, data, sepsis_risk, latency_ms)
return PredictionResponse(
patient_id=data.patient_id,
sepsis_risk=sepsis_risk,
risk_category=risk_category,
confidence=confidence,
model_version=model_metadata.get('version', '1.0.0'),
timestamp=datetime.utcnow().isoformat(),
prediction_id=f"pred_{prediction_id}"
)
except ValueError as e:
logger.error(f"Validation error: {e}")
raise HTTPException(status_code=400, detail=f"Invalid input: {str(e)}")
except Exception as e:
logger.error(f"Prediction error: {e}", exc_info=True)
raise HTTPException(status_code=500, detail=f"Prediction failed: {str(e)}")
@app.post("/api/predict/batch")
async def predict_batch(
data: List[PatientData],
api_key: str = Depends(verify_api_key)
):
"""
Batch prediction endpoint for multiple patients
**Limits:** Maximum 100 patients per request
"""
if len(data) > 100:
raise HTTPException(
status_code=400,
detail="Maximum 100 patients per batch request"
)
try:
# Convert to DataFrame
patient_ids = [d.patient_id for d in data]
input_df = pd.DataFrame([{
'heart_rate': d.heart_rate,
'respiratory_rate': d.respiratory_rate,
'temperature': d.temperature,
'systolic_bp': d.systolic_bp,
'white_blood_cell': d.white_blood_cell,
'lactate': d.lactate,
'age': d.age
} for d in data])
# Make predictions
predictions = model.predict(input_df)
# Format results
results = []
for patient_id, pred in zip(patient_ids, predictions):
sepsis_risk = float(pred)
risk_category = (
"Low" if sepsis_risk < 0.3 else
"Medium" if sepsis_risk < 0.7 else
"High"
)
confidence = min(1.0, max(abs(sepsis_risk - 0.5) * 2, 0.5))
prediction_id = hashlib.sha256(
f"{patient_id}{datetime.utcnow().isoformat()}".encode()
).hexdigest()[:16]
results.append({
"patient_id": patient_id,
"sepsis_risk": sepsis_risk,
"risk_category": risk_category,
"confidence": confidence,
"model_version": model_metadata.get('version', '1.0.0'),
"timestamp": datetime.utcnow().isoformat(),
"prediction_id": f"pred_{prediction_id}"
})
logger.info(f"Batch prediction completed: {len(results)} patients")
return {"predictions": results, "count": len(results)}
except Exception as e:
logger.error(f"Batch prediction error: {e}", exc_info=True)
raise HTTPException(status_code=500, detail=f"Batch prediction failed: {str(e)}")
@app.get("/api/model/info")
async def model_info(api_key: str = Depends(verify_api_key)):
"""Get information about the deployed model"""
return {
"model_name": "sepsis_predictor",
"version": model_metadata.get('version'),
"run_id": model_metadata.get('run_id'),
"deployed_at": model_metadata.get('created_at'),
"features": [
"heart_rate", "respiratory_rate", "temperature",
"systolic_bp", "white_blood_cell", "lactate", "age"
],
"performance": {
"validation_auc": model_metadata.get('tags', {}).get('val_auc'),
"validation_sensitivity": model_metadata.get('tags', {}).get('val_sensitivity'),
"validation_specificity": model_metadata.get('tags', {}).get('val_specificity')
}
}
# Run with: uvicorn api:app --host 0.0.0.0 --port 8000 --workers 4Testing the API:
import requests
import json
# Test single prediction
url = "http://localhost:8000/api/predict"
headers = {"X-API-Key": "demo_key_123", "Content-Type": "application/json"}
patient_data = {
"patient_id": "ICU-2024-001",
"heart_rate": 110.5,
"respiratory_rate": 24.0,
"temperature": 38.5,
"systolic_bp": 95.0,
"white_blood_cell": 15.2,
"lactate": 2.8,
"age": 67
}
response = requests.post(url, headers=headers, json=patient_data)
print(f"Status: {response.status_code}")
print(f"Response: {json.dumps(response.json(), indent=2)}")
# Test batch prediction
batch_url = "http://localhost:8000/api/predict/batch"
batch_data = [patient_data, {**patient_data, "patient_id": "ICU-2024-002", "lactate": 1.5}]
batch_response = requests.post(batch_url, headers=headers, json=batch_data)
print(f"\nBatch predictions: {len(batch_response.json()['predictions'])}")Option 2: Batch Predictions
Use case: Non-real-time predictions for entire cohorts, scheduled risk stratification
import pandas as pd
import mlflow.pyfunc
from datetime import datetime
import logging
from typing import Optional
import argparse
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)
def batch_predict(
input_csv: str,
output_csv: str,
model_name: str = "sepsis_predictor",
stage: str = "Production",
chunk_size: int = 1000
):
"""
Run batch predictions on a CSV file
Args:
input_csv: Path to input CSV with patient data
output_csv: Path to save predictions
model_name: Name of registered model in MLflow
stage: Model stage (Production, Staging, None)
chunk_size: Process data in chunks for memory efficiency
"""
# Load model
model_uri = f"models:/{model_name}/{stage}"
logger.info(f"Loading model from {model_uri}")
model = mlflow.pyfunc.load_model(model_uri)
# Load data in chunks for large files
logger.info(f"Reading input data from {input_csv}")
chunks = []
total_rows = 0
for chunk in pd.read_csv(input_csv, chunksize=chunk_size):
total_rows += len(chunk)
logger.info(f"Loaded {total_rows} rows...")
chunks.append(chunk)
df = pd.concat(chunks, ignore_index=True)
logger.info(f"[OK] Loaded {len(df)} patients for batch prediction")
# Store patient IDs
patient_ids = df['patient_id'].values
# Prepare features
feature_cols = [
'heart_rate', 'respiratory_rate', 'temperature',
'systolic_bp', 'white_blood_cell', 'lactate', 'age'
]
# Validate required columns
missing_cols = [col for col in feature_cols if col not in df.columns]
if missing_cols:
raise ValueError(f"Missing required columns: {missing_cols}")
X = df[feature_cols]
# Make predictions in chunks
logger.info("Making predictions...")
all_predictions = []
for i in range(0, len(X), chunk_size):
chunk_X = X.iloc[i:i+chunk_size]
chunk_preds = model.predict(chunk_X)
all_predictions.extend(chunk_preds)
logger.info(f"Processed {min(i+chunk_size, len(X))}/{len(X)} predictions")
predictions = np.array(all_predictions)
# Add predictions to dataframe
df['sepsis_risk'] = predictions
df['risk_category'] = pd.cut(
predictions,
bins=[0, 0.3, 0.7, 1.0],
labels=['Low', 'Medium', 'High']
)
df['prediction_timestamp'] = datetime.utcnow().isoformat()
df['model_version'] = stage
# Save results
logger.info(f"Saving predictions to {output_csv}")
df.to_csv(output_csv, index=False)
# Log summary statistics
summary = df['risk_category'].value_counts().to_dict()
logger.info(f"[OK] Batch prediction complete!")
logger.info(f"Risk distribution: {summary}")
logger.info(f"Mean risk: {predictions.mean():.3f}")
logger.info(f"Median risk: {np.median(predictions):.3f}")
# Identify high-risk patients
high_risk = df[df['sepsis_risk'] >= 0.7]
logger.info(f"[WARNING] {len(high_risk)} high-risk patients identified")
if len(high_risk) > 0:
logger.info(f"High-risk patient IDs (first 10): {high_risk['patient_id'].head(10).tolist()}")
return df
def main():
parser = argparse.ArgumentParser(description='Batch sepsis risk prediction')
parser.add_argument('--input', required=True, help='Input CSV file')
parser.add_argument('--output', required=True, help='Output CSV file')
parser.add_argument('--model', default='sepsis_predictor', help='Model name')
parser.add_argument('--stage', default='Production', help='Model stage')
parser.add_argument('--chunk-size', type=int, default=1000, help='Chunk size')
args = parser.parse_args()
batch_predict(
input_csv=args.input,
output_csv=args.output,
model_name=args.model,
stage=args.stage,
chunk_size=args.chunk_size
)
if __name__ == "__main__":
main()
# Usage:
# python batch_predict.py \
# --input data/icu_patients_2024-03.csv \
# --output predictions/sepsis_risk_2024-03.csv \
# --model sepsis_predictor \
# --stage ProductionOption 3: Embedded in Application
Use case: Model runs inside existing application (e.g., EHR plugin, mobile app)
import mlflow.pyfunc
from typing import Dict, List, Optional
import logging
class EmbeddedSepsisPredictor:
"""
Embedded sepsis predictor for integration into EHR systems
Designed to run as part of EHR workflow without external API calls
"""
def __init__(self, model_path: str):
"""
Initialize predictor
Args:
model_path: Local path to model or MLflow model URI
"""
self.model = mlflow.pyfunc.load_model(model_path)
self.feature_names = [
'heart_rate', 'respiratory_rate', 'temperature',
'systolic_bp', 'white_blood_cell', 'lactate', 'age'
]
self.logger = logging.getLogger(__name__)
def predict_from_ehr(self, patient_record: Dict) -> Dict:
"""
Extract features from EHR record and make prediction
Args:
patient_record: Dict containing patient data (FHIR or custom format)
Returns:
Dict with prediction results and interpretable output
"""
try:
# Extract features from EHR record
features = self._extract_features(patient_record)
# Validate features
if not self._validate_features(features):
return {
'success': False,
'error': 'Insufficient data for prediction',
'missing_features': [f for f in self.feature_names if f not in features]
}
# Make prediction
import pandas as pd
features_df = pd.DataFrame([features])
risk = self.model.predict(features_df)[0]
# Generate explanation
explanation = self._generate_explanation(features, risk)
return {
'success': True,
'patient_id': patient_record.get('patient_id'),
'sepsis_risk': float(risk),
'risk_category': self._categorize_risk(risk),
'features_used': features,
'explanation': explanation,
'model_version': '1.2.0',
'timestamp': datetime.now().isoformat()
}
except Exception as e:
self.logger.error(f"Prediction error: {e}", exc_info=True)
return {
'success': False,
'error': str(e)
}
def _extract_features(self, record: Dict) -> Dict:
"""
Extract features from EHR record
Supports both FHIR format and custom EHR formats
"""
features = {}
# Extract from vitals (assuming EHR provides recent vitals)
vitals = record.get('vitals', {})
features['heart_rate'] = vitals.get('heart_rate')
features['respiratory_rate'] = vitals.get('respiratory_rate')
features['temperature'] = vitals.get('temperature')
features['systolic_bp'] = vitals.get('systolic_bp')
# Extract from labs (most recent values)
labs = record.get('labs', {})
features['white_blood_cell'] = labs.get('wbc')
features['lactate'] = labs.get('lactate')
# Calculate age from birth date
if 'birth_date' in record:
from datetime import datetime
from dateutil.parser import parse
birth_date = parse(record['birth_date'])
features['age'] = (datetime.now() - birth_date).days // 365
return features
def _validate_features(self, features: Dict) -> bool:
"""Check if minimum required features are present"""
required = ['heart_rate', 'temperature', 'age']
return all(features.get(f) is not None for f in required)
def _categorize_risk(self, risk: float) -> str:
"""Categorize risk level"""
if risk < 0.3:
return "Low"
elif risk < 0.7:
return "Medium"
else:
return "High"
def _generate_explanation(self, features: Dict, risk: float) -> str:
"""
Generate human-readable explanation of prediction
In production: Use SHAP, LIME, or similar for proper explanations
"""
contributors = []
# Simplified rule-based explanation
if features.get('lactate', 0) > 2.0:
contributors.append(f"elevated lactate ({features['lactate']:.1f} mmol/L)")
if features.get('heart_rate', 0) > 100:
contributors.append(f"tachycardia (HR {features['heart_rate']:.0f} bpm)")
if features.get('temperature', 37) > 38.0:
contributors.append(f"fever ({features['temperature']:.1f}°C)")
if features.get('white_blood_cell', 0) > 12:
contributors.append(f"leukocytosis (WBC {features['white_blood_cell']:.1f} K/μL)")
if features.get('systolic_bp', 120) < 100:
contributors.append(f"hypotension (SBP {features['systolic_bp']:.0f} mmHg)")
if contributors:
return f"Risk elevated due to: {', '.join(contributors)}"
else:
return "Vital signs and labs within normal ranges"
# Example usage in EHR system
if __name__ == "__main__":
# Initialize predictor (done once at application startup)
predictor = EmbeddedSepsisPredictor(model_path="models:/sepsis_predictor/Production")
# Example patient record from EHR
patient_record = {
'patient_id': 'MRN12345',
'birth_date': '1957-03-15',
'vitals': {
'heart_rate': 112,
'respiratory_rate': 24,
'temperature': 38.5,
'systolic_bp': 92
},
'labs': {
'wbc': 16.2,
'lactate': 3.1
}
}
# Make prediction
result = predictor.predict_from_ehr(patient_record)
if result['success']:
print(f"Patient {result['patient_id']}:")
print(f" Sepsis Risk: {result['sepsis_risk']:.1%}")
print(f" Category: {result['risk_category']}")
print(f" Explanation: {result['explanation']}")
else:
print(f"Prediction failed: {result['error']}")Option 4: Edge Deployment
Use case: Mobile devices, resource-constrained environments, offline operation
Howard et al., 2017, arXiv - MobileNets: Efficient CNNs for mobile vision
import tensorflow as tf
import coremltools as ct
import numpy as np
def convert_to_tflite(keras_model_path: str, output_path: str, quantize: bool = True):
"""
Convert Keras model to TensorFlow Lite for mobile deployment
Args:
keras_model_path: Path to saved Keras model
output_path: Path to save TFLite model
quantize: Whether to apply quantization for size reduction
"""
# Load Keras model
model = tf.keras.models.load_model(keras_model_path)
# Convert to TFLite
converter = tf.lite.TFLiteConverter.from_keras_model(model)
if quantize:
# Apply post-training quantization
converter.optimizations = [tf.lite.Optimize.DEFAULT]
converter.target_spec.supported_types = [tf.float16]
# For even smaller size, use int8 quantization
# converter.target_spec.supported_ops = [tf.lite.OpsSet.TFLITE_BUILTINS_INT8]
# converter.inference_input_type = tf.uint8
# converter.inference_output_type = tf.uint8
tflite_model = converter.convert()
# Save
with open(output_path, 'wb') as f:
f.write(tflite_model)
# Print model size
import os
size_mb = os.path.getsize(output_path) / (1024 * 1024)
print(f"[OK] TFLite model saved to {output_path}")
print(f"Model size: {size_mb:.2f} MB")
# Benchmark inference speed
interpreter = tf.lite.Interpreter(model_path=output_path)
interpreter.allocate_tensors()
input_details = interpreter.get_input_details()
output_details = interpreter.get_output_details()
# Test inference
test_input = np.random.randn(1, 7).astype(np.float32)
import time
start = time.time()
for _ in range(100):
interpreter.set_tensor(input_details[0]['index'], test_input)
interpreter.invoke()
_ = interpreter.get_tensor(output_details[0]['index'])
end = time.time()
avg_latency_ms = (end - start) / 100 * 1000
print(f"Average inference latency: {avg_latency_ms:.2f} ms")
return output_path
def convert_to_coreml(keras_model_path: str, output_path: str):
"""
Convert Keras model to Core ML for iOS deployment
Args:
keras_model_path: Path to saved Keras model
output_path: Path to save Core ML model (.mlmodel)
"""
# Load model
model = tf.keras.models.load_model(keras_model_path)
# Convert to Core ML
coreml_model = ct.convert(
model,
inputs=[ct.TensorType(name="input", shape=(1, 7))],
convert_to="mlprogram", # Use ML Program (newer format)
minimum_deployment_target=ct.target.iOS15
)
# Add metadata
coreml_model.author = "Hospital ML Team"
coreml_model.short_description = "Sepsis risk prediction model"
coreml_model.version = "1.2.0"
# Add input/output descriptions
coreml_model.input_description['input'] = (
"7 features: heart rate, respiratory rate, temperature, "
"systolic BP, WBC, lactate, age"
)
coreml_model.output_description['output'] = "Sepsis risk probability (0-1)"
# Save
coreml_model.save(output_path)
import os
size_mb = os.path.getsize(output_path) / (1024 * 1024)
print(f"[OK] Core ML model saved to {output_path}")
print(f"Model size: {size_mb:.2f} MB")
return output_path
# Example: iOS Swift code to use Core ML model
ios_swift_code = '''
import CoreML
class SepsisPredictor {
let model: SepsisRiskModel // Auto-generated from .mlmodel
init() {
do {
model = try SepsisRiskModel(configuration: MLModelConfiguration())
} catch {
fatalError("Failed to load model: \\(error)")
}
}
func predict(heartRate: Double, respiratoryRate: Double, temperature: Double,
systolicBP: Double, wbc: Double, lactate: Double, age: Double) -> Double? {
// Prepare input
let input = try? SepsisRiskModelInput(
input: [heartRate, respiratoryRate, temperature, systolicBP, wbc, lactate, age]
)
guard let input = input else { return nil }
// Make prediction
guard let output = try? model.prediction(input: input) else { return nil }
// Extract probability
return output.output[0]
}
}
// Usage
let predictor = SepsisPredictor()
let risk = predictor.predict(
heartRate: 110,
respiratoryRate: 24,
temperature: 38.5,
systolicBP: 95,
wbc: 15.2,
lactate: 2.8,
age: 67
)
if let risk = risk {
print("Sepsis risk: \\(risk * 100)%")
}
'''
# Convert models
if __name__ == "__main__":
# Convert to TensorFlow Lite
convert_to_tflite(
keras_model_path="models/sepsis_model.h5",
output_path="models/sepsis_model.tflite",
quantize=True
)
# Convert to Core ML
convert_to_coreml(
keras_model_path="models/sepsis_model.h5",
output_path="models/SepsisRiskModel.mlmodel"
)Deployment Patterns
Blue-Green Deployment
Concept: Run two identical production environments (“blue” and “green”). Deploy new version to inactive environment, test, then switch traffic.
Fowler, 2010 - BlueGreenDeployment
Advantages: - Zero downtime - Instant rollback (switch back to blue) - Full testing in production environment before cutover
Kubernetes implementation:
# deployment-blue.yaml
apiVersion: apps/v1
kind: Deployment
metadata:
name: sepsis-predictor-blue
namespace: production
spec:
replicas: 3
selector:
matchLabels:
app: sepsis-predictor
version: blue
template:
metadata:
labels:
app: sepsis-predictor
version: blue
spec:
containers:
- name: api
image: sepsis-predictor:v1.0.0
ports:
- containerPort: 8000
env:
- name: MODEL_URI
value: "models:/sepsis_predictor/Production"
resources:
requests:
memory: "512Mi"
cpu: "500m"
limits:
memory: "1Gi"
cpu: "1000m"
livenessProbe:
httpGet:
path: /health
port: 8000
initialDelaySeconds: 30
periodSeconds: 10
readinessProbe:
httpGet:
path: /health
port: 8000
initialDelaySeconds: 5
periodSeconds: 5
---
# deployment-green.yaml
apiVersion: apps/v1
kind: Deployment
metadata:
name: sepsis-predictor-green
namespace: production
spec:
replicas: 3
selector:
matchLabels:
app: sepsis-predictor
version: green
template:
metadata:
labels:
app: sepsis-predictor
version: green
spec:
containers:
- name: api
image: sepsis-predictor:v1.1.0 # New version
ports:
- containerPort: 8000
env:
- name: MODEL_URI
value: "models:/sepsis_predictor/Production"
resources:
requests:
memory: "512Mi"
cpu: "500m"
limits:
memory: "1Gi"
cpu: "1000m"
livenessProbe:
httpGet:
path: /health
port: 8000
initialDelaySeconds: 30
periodSeconds: 10
readinessProbe:
httpGet:
path: /health
port: 8000
initialDelaySeconds: 5
periodSeconds: 5
---
# service.yaml
apiVersion: v1
kind: Service
metadata:
name: sepsis-predictor
namespace: production
spec:
selector:
app: sepsis-predictor
version: blue # Initially points to blue
ports:
- protocol: TCP
port: 80
targetPort: 8000
type: LoadBalancerDeployment script:
#!/bin/bash
# blue_green_deploy.sh
set -e
# Colors for output
RED='\033[0;31m'
GREEN='\033[0;32m'
YELLOW='\033[1;33m'
NC='\033[0m' # No Color
echo -e "${GREEN}Starting blue-green deployment...${NC}"
# Deploy green version
echo -e "${YELLOW}Deploying green version...${NC}"
kubectl apply -f k8s/deployment-green.yaml
# Wait for green deployment to be ready
echo -e "${YELLOW}Waiting for green deployment to be ready...${NC}"
kubectl wait --for=condition=available --timeout=300s \
deployment/sepsis-predictor-green -n production
if [ $? -ne 0 ]; then
echo -e "${RED}Green deployment failed to become ready${NC}"
exit 1
fi
# Run smoke tests against green
echo -e "${YELLOW}Running smoke tests against green...${NC}"
GREEN_IP=$(kubectl get svc sepsis-predictor-green -n production -o jsonpath='{.status.loadBalancer.ingress[0].ip}')
python tests/smoke_test.py --endpoint "http://${GREEN_IP}" --timeout 60
if [ $? -eq 0 ]; then
echo -e "${GREEN}Smoke tests passed!${NC}"
# Switch traffic from blue to green
echo -e "${YELLOW}Switching traffic to green...${NC}"
kubectl patch service sepsis-predictor \
-p '{"spec":{"selector":{"version":"green"}}}' \
-n production
echo -e "${GREEN}Traffic switched to green${NC}"
# Monitor for 10 minutes
echo -e "${YELLOW}Monitoring green deployment for 10 minutes...${NC}"
sleep 600
# Check if errors occurred
ERROR_RATE=$(kubectl logs deployment/sepsis-predictor-green -n production | \
grep -i error | wc -l)
if [ $ERROR_RATE -gt 10 ]; then
echo -e "${RED}High error rate detected, rolling back to blue${NC}"
kubectl patch service sepsis-predictor \
-p '{"spec":{"selector":{"version":"blue"}}}' \
-n production
exit 1
fi
# If successful, scale down blue
echo -e "${GREEN}Deployment successful, scaling down blue${NC}"
kubectl scale deployment sepsis-predictor-blue --replicas=0 -n production
echo -e "${GREEN}Blue-green deployment complete!${NC}"
else
echo -e "${RED}Smoke tests failed, keeping blue active${NC}"
kubectl scale deployment sepsis-predictor-green --replicas=0 -n production
exit 1
fiCanary Deployment
Concept: Gradually shift traffic from old to new version, monitoring closely at each step.
Humble & Farley, 2010, Continuous Delivery
Advantages: - Lower risk than immediate full cutover - Real-world validation with subset of users - Easy to halt if issues detected
Istio VirtualService for canary:
apiVersion: networking.istio.io/v1alpha3
kind: VirtualService
metadata:
name: sepsis-predictor-canary
namespace: production
spec:
hosts:
- sepsis-predictor.production.svc.cluster.local
http:
# Route 10% of traffic to v2 (canary)
- match:
- headers:
x-canary-user:
exact: "true"
route:
- destination:
host: sepsis-predictor
subset: v2
- route:
- destination:
host: sepsis-predictor
subset: v1
weight: 90
- destination:
host: sepsis-predictor
subset: v2
weight: 10 # Start with 10% trafficProgressive rollout script:
import subprocess
import time
import requests
from typing import Dict
def get_error_rate(version: str) -> float:
"""Query Prometheus for error rate of specific version"""
query = f'rate(http_requests_total{{version="{version}",status=~"5.."}}[5m])'
response = requests.get(
'http://prometheus:9090/api/v1/query',
params={'query': query}
)
result = response.json()['data']['result']
if result:
return float(result[0]['value'][1])
return 0.0
def get_latency_p95(version: str) -> float:
"""Get 95th percentile latency"""
query = f'histogram_quantile(0.95, rate(http_request_duration_seconds_bucket{{version="{version}"}}[5m]))'
response = requests.get(
'http://prometheus:9090/api/v1/query',
params={'query': query}
)
result = response.json()['data']['result']
if result:
return float(result[0]['value'][1])
return 0.0
def update_traffic_split(v1_weight: int, v2_weight: int):
"""Update Istio VirtualService traffic weights"""
yaml = f"""
apiVersion: networking.istio.io/v1alpha3
kind: VirtualService
metadata:
name: sepsis-predictor-canary
namespace: production
spec:
hosts:
- sepsis-predictor.production.svc.cluster.local
http:
- route:
- destination:
host: sepsis-predictor
subset: v1
weight: {v1_weight}
- destination:
host: sepsis-predictor
subset: v2
weight: {v2_weight}
"""
with open('/tmp/vs.yaml', 'w') as f:
f.write(yaml)
subprocess.run(['kubectl', 'apply', '-f', '/tmp/vs.yaml'], check=True)
print(f"[OK] Updated traffic split: v1={v1_weight}%, v2={v2_weight}%")
def canary_deployment():
"""
Progressive canary deployment with automated monitoring
Progression: 10% → 25% → 50% → 75% → 100%
"""
stages = [
(90, 10), # 10% canary
(75, 25), # 25% canary
(50, 50), # 50% canary
(25, 75), # 75% canary
(0, 100) # 100% canary (full rollout)
]
# Error rate and latency thresholds
MAX_ERROR_RATE = 0.05 # 5%
MAX_LATENCY_P95 = 0.5 # 500ms
for i, (v1_weight, v2_weight) in enumerate(stages):
print(f"\n{'='*60}")
print(f"Stage {i+1}/{len(stages)}: Shifting to {v2_weight}% canary")
print(f"{'='*60}")
# Update traffic split
update_traffic_split(v1_weight, v2_weight)
# Monitor for 10 minutes
monitoring_duration = 600 # 10 minutes
check_interval = 30 # Check every 30 seconds
for elapsed in range(0, monitoring_duration, check_interval):
time.sleep(check_interval)
# Get metrics for canary (v2)
v2_error_rate = get_error_rate('v2')
v2_latency_p95 = get_latency_p95('v2')
# Get metrics for stable (v1) for comparison
v1_error_rate = get_error_rate('v1')
v1_latency_p95 = get_latency_p95('v1')
print(f"\nElapsed: {elapsed}/{monitoring_duration}s")
print(f"V2 (canary) - Error rate: {v2_error_rate:.2%}, P95 latency: {v2_latency_p95*1000:.0f}ms")
print(f"V1 (stable) - Error rate: {v1_error_rate:.2%}, P95 latency: {v1_latency_p95*1000:.0f}ms")
# Check if canary exceeds thresholds
if v2_error_rate > MAX_ERROR_RATE:
print(f"\n[FAILED] Canary error rate {v2_error_rate:.2%} exceeds threshold {MAX_ERROR_RATE:.2%}")
print("Rolling back to v1...")
update_traffic_split(100, 0)
return False
if v2_latency_p95 > MAX_LATENCY_P95:
print(f"\n[FAILED] Canary latency {v2_latency_p95*1000:.0f}ms exceeds threshold {MAX_LATENCY_P95*1000:.0f}ms")
print("Rolling back to v1...")
update_traffic_split(100, 0)
return False
# Check if canary significantly worse than stable
if v2_error_rate > v1_error_rate * 1.5:
print(f"\n[WARNING] Canary error rate 50% higher than stable")
print("Rolling back to v1...")
update_traffic_split(100, 0)
return False
print(f"\n[OK] Stage {i+1} monitoring complete - metrics within acceptable range")
print("\n" + "="*60)
print("Canary deployment successful! v2 now receiving 100% of traffic")
print("="*60)
return True
if __name__ == "__main__":
success = canary_deployment()
if not success:
print("\n[FAILED] Canary deployment failed and was rolled back")
exit(1)For additional deployment patterns, see Richardson, 2018, Microservices Patterns.
Production Monitoring
What to Monitor
Three critical categories:
- Model Performance - Is the model still accurate?
- Data Quality - Are inputs valid and within expected ranges?
- System Health - Is the service responsive and available?
Breck et al., 2017, IEEE Big Data - “The ML Test Score: A Rubric for ML Production Readiness”
Performance Monitoring
Metrics to Track with Prometheus
Prometheus documentation - Industry-standard monitoring system
from prometheus_client import Counter, Histogram, Gauge, Summary
import prometheus_client
from fastapi import FastAPI
import time
import numpy as np
app = FastAPI()
# Define metrics
prediction_counter = Counter(
'sepsis_predictions_total',
'Total number of sepsis predictions',
['risk_category', 'model_version']
)
prediction_latency = Histogram(
'sepsis_prediction_latency_seconds',
'Prediction latency in seconds',
buckets=[0.001, 0.005, 0.01, 0.025, 0.05, 0.1, 0.25, 0.5, 1.0, 2.5, 5.0]
)
active_predictions = Gauge(
'sepsis_active_predictions',
'Number of predictions currently being processed'
)
model_confidence = Histogram(
'sepsis_model_confidence',
'Distribution of model confidence scores',
buckets=[0.0, 0.1, 0.2, 0.3, 0.4, 0.5, 0.6, 0.7, 0.8, 0.9, 1.0]
)
input_feature_distribution = Histogram(
'sepsis_input_feature_value',
'Distribution of input feature values',
['feature_name'],
buckets=list(np.percentile(range(0, 100), np.arange(0, 101, 10)))
)
model_errors = Counter(
'sepsis_prediction_errors_total',
'Total number of prediction errors',
['error_type']
)
# Actual outcomes (when available) for performance tracking
actual_outcomes = Counter(
'sepsis_actual_outcomes_total',
'Actual sepsis outcomes',
['predicted_category', 'actual_outcome']
)
@app.post("/predict")
@active_predictions.track_inprogress()
def predict(data: PatientData):
start_time = time.time()
try:
# Prepare input
input_df = pd.DataFrame([data.dict()])
# Log input feature distributions
for feature, value in data.dict().items():
if isinstance(value, (int, float)):
input_feature_distribution.labels(feature_name=feature).observe(value)
# Make prediction
prediction = model.predict(input_df)[0]
sepsis_risk = float(prediction)
# Categorize risk
if sepsis_risk < 0.3:
risk_category = "Low"
elif sepsis_risk < 0.7:
risk_category = "Medium"
else:
risk_category = "High"
# Record metrics
prediction_counter.labels(
risk_category=risk_category,
model_version="1.2.0"
).inc()
# Model confidence (distance from decision boundary)
confidence = min(1.0, max(abs(sepsis_risk - 0.5) * 2, 0.5))
model_confidence.observe(confidence)
# Record latency
latency = time.time() - start_time
prediction_latency.observe(latency)
logger.info(
f"Prediction: patient={data.patient_id}, risk={sepsis_risk:.3f}, "
f"category={risk_category}, latency={latency*1000:.1f}ms"
)
return {
"sepsis_risk": sepsis_risk,
"risk_category": risk_category,
"confidence": confidence
}
except ValueError as e:
model_errors.labels(error_type="validation_error").inc()
raise HTTPException(status_code=400, detail=str(e))
except Exception as e:
model_errors.labels(error_type="prediction_error").inc()
logger.error(f"Prediction error: {e}", exc_info=True)
raise HTTPException(status_code=500, detail=str(e))
@app.post("/feedback")
def record_outcome(patient_id: str, actual_sepsis: bool, predicted_risk: float):
"""
Record actual outcome for model performance monitoring
Called after clinical confirmation of sepsis status
"""
predicted_category = (
"Low" if predicted_risk < 0.3 else
"Medium" if predicted_risk < 0.7 else
"High"
)
actual_outcomes.labels(
predicted_category=predicted_category,
actual_outcome="sepsis" if actual_sepsis else "no_sepsis"
).inc()
logger.info(f"Outcome recorded: patient={patient_id}, actual={actual_sepsis}, predicted_risk={predicted_risk:.3f}")
return {"status": "recorded"}
@app.get("/metrics")
def metrics():
"""Prometheus metrics endpoint"""
return Response(
prometheus_client.generate_latest(),
media_type="text/plain"
)Prometheus configuration:
# prometheus.yml
global:
scrape_interval: 15s
evaluation_interval: 15s
scrape_configs:
- job_name: 'sepsis-predictor'
static_configs:
- targets: ['sepsis-predictor:8000']
metrics_path: '/metrics'
# Alert rules
rule_files:
- 'alerts.yml'
alerting:
alertmanagers:
- static_configs:
- targets: ['alertmanager:9093']Alert rules:
# alerts.yml
groups:
- name: sepsis_predictor_alerts
interval: 30s
rules:
# High error rate
- alert: HighErrorRate
expr: rate(sepsis_prediction_errors_total[5m]) > 0.05
for: 5m
labels:
severity: critical
annotations:
summary: "High prediction error rate"
description: "Error rate is {{ $value }} errors/sec"
# High latency
- alert: HighLatency
expr: histogram_quantile(0.95, rate(sepsis_prediction_latency_seconds_bucket[5m])) > 1.0
for: 5m
labels:
severity: warning
annotations:
summary: "High prediction latency"
description: "P95 latency is {{ $value }}s"
# Model performance degradation
- alert: ModelPerformanceDegradation
expr: |
sum(rate(sepsis_actual_outcomes_total{predicted_category="High",actual_outcome="no_sepsis"}[1h])) /
sum(rate(sepsis_actual_outcomes_total{predicted_category="High"}[1h])) > 0.5
for: 1h
labels:
severity: warning
annotations:
summary: "High false positive rate detected"
description: "FPR for high-risk predictions is {{ $value | humanizePercentage }}"
# Low prediction volume (possible system issue)
- alert: LowPredictionVolume
expr: rate(sepsis_predictions_total[5m]) < 0.1
for: 10m
labels:
severity: warning
annotations:
summary: "Unusually low prediction volume"
description: "Only {{ $value }} predictions/sec"Data Drift Detection
Monitor distribution shifts in input features over time.
Rabanser et al., 2019, NeurIPS - “Failing Loudly: An Empirical Study of Methods for Detecting Dataset Shift”
import numpy as np
import pandas as pd
from scipy import stats
from typing import Dict, List, Tuple
from dataclasses import dataclass
from datetime import datetime, timedelta
@dataclass
class DriftResult:
feature: str
ks_statistic: float
p_value: float
psi: float
drift_detected: bool
severity: str # 'none', 'moderate', 'severe'
class DataDriftDetector:
"""
Detect data drift using multiple statistical tests
Methods:
- Kolmogorov-Smirnov test for distribution shift
- Population Stability Index (PSI) for drift magnitude
"""
def __init__(self, reference_data: pd.DataFrame, threshold: float = 0.05):
"""
Args:
reference_data: Training/baseline data distribution
threshold: P-value threshold for KS test (default 0.05)
"""
self.reference_data = reference_data
self.threshold = threshold
self.reference_stats = self._calculate_stats(reference_data)
self.drift_history = []
def _calculate_stats(self, df: pd.DataFrame) -> Dict:
"""Calculate distribution statistics for reference"""
stats_dict = {}
for col in df.select_dtypes(include=[np.number]).columns:
stats_dict[col] = {
'mean': df[col].mean(),
'std': df[col].std(),
'median': df[col].median(),
'min': df[col].min(),
'max': df[col].max(),
'quartiles': df[col].quantile([0.25, 0.5, 0.75]).to_dict(),
'skew': df[col].skew(),
'kurtosis': df[col].kurtosis()
}
return stats_dict
def detect_drift(self, current_data: pd.DataFrame) -> List[DriftResult]:
"""
Detect drift using KS test and PSI
Returns:
List of DriftResult objects, one per feature
"""
results = []
for col in self.reference_data.select_dtypes(include=[np.number]).columns:
if col not in current_data.columns:
continue
# Remove NaN values
ref_values = self.reference_data[col].dropna()
cur_values = current_data[col].dropna()
if len(cur_values) < 30: # Insufficient data
continue
# Kolmogorov-Smirnov test
ks_statistic, p_value = stats.ks_2samp(ref_values, cur_values)
# Population Stability Index
psi = self._calculate_psi(ref_values, cur_values)
# Determine drift severity
drift_detected = p_value < self.threshold
if psi < 0.1:
severity = 'none'
elif psi < 0.2:
severity = 'moderate'
else:
severity = 'severe'
result = DriftResult(
feature=col,
ks_statistic=float(ks_statistic),
p_value=float(p_value),
psi=float(psi),
drift_detected=drift_detected,
severity=severity
)
results.append(result)
# Store in history
self.drift_history.append({
'timestamp': datetime.now(),
'results': results
})
return results
def _calculate_psi(self, expected: pd.Series, actual: pd.Series, bins: int = 10) -> float:
"""
Calculate Population Stability Index (PSI)
PSI interpretation:
- < 0.1: No significant change
- 0.1-0.2: Moderate change (investigate)
- > 0.2: Significant change (retrain model)
Formula: PSI = Σ (actual% - expected%) × ln(actual% / expected%)
"""
# Create bins based on expected distribution percentiles
breakpoints = np.percentile(expected, np.linspace(0, 100, bins + 1))
breakpoints = np.unique(breakpoints) # Remove duplicates
if len(breakpoints) < 3:
# Not enough unique values for binning
return 0.0
# Calculate distributions
expected_dist = np.histogram(expected, bins=breakpoints)[0] / len(expected)
actual_dist = np.histogram(actual, bins=breakpoints)[0] / len(actual)
# Add small constant to avoid log(0)
expected_dist = expected_dist + 0.0001
actual_dist = actual_dist + 0.0001
# Calculate PSI
psi = np.sum((actual_dist - expected_dist) * np.log(actual_dist / expected_dist))
return psi
def get_drift_report(self, results: List[DriftResult]) -> str:
"""Generate human-readable drift report"""
report = "="*60 + "\n"
report += "DATA DRIFT DETECTION REPORT\n"
report += f"Timestamp: {datetime.now().isoformat()}\n"
report += "="*60 + "\n\n"
# Count drifted features
drifted = [r for r in results if r.drift_detected]
report += f"Features analyzed: {len(results)}\n"
report += f"Features with drift: {len(drifted)}\n\n"
if len(drifted) == 0:
report += "[OK] No significant drift detected\n"
else:
report += "[WARNING] DRIFT DETECTED:\n\n"
for result in sorted(drifted, key=lambda x: x.psi, reverse=True):
report += f"Feature: {result.feature}\n"
report += f" KS Statistic: {result.ks_statistic:.4f}\n"
report += f" P-value: {result.p_value:.4f}\n"
report += f" PSI: {result.psi:.4f}\n"
report += f" Severity: {result.severity.upper()}\n"
# Get reference vs current stats
ref_mean = self.reference_stats[result.feature]['mean']
cur_mean = self.reference_data[result.feature].mean() # This should be current_data
report += f" Reference mean: {ref_mean:.2f}\n"
report += f" Current mean: {cur_mean:.2f}\n"
report += f" Change: {((cur_mean - ref_mean) / ref_mean * 100):+.1f}%\n\n"
return report
def plot_drift(self, results: List[DriftResult], output_path: str = 'drift_report.html'):
"""Generate interactive drift visualization"""
import plotly.graph_objects as go
from plotly.subplots import make_subplots
# Create subplots
fig = make_subplots(
rows=2, cols=1,
subplot_titles=('PSI Scores by Feature', 'P-values by Feature')
)
# Sort by PSI
results_sorted = sorted(results, key=lambda x: x.psi, reverse=True)
features = [r.feature for r in results_sorted]
psi_values = [r.psi for r in results_sorted]
p_values = [r.p_value for r in results_sorted]
# PSI bar chart
colors_psi = [
'green' if psi < 0.1 else 'orange' if psi < 0.2 else 'red'
for psi in psi_values
]
fig.add_trace(
go.Bar(x=features, y=psi_values, marker_color=colors_psi, name='PSI'),
row=1, col=1
)
# Add PSI threshold lines
fig.add_hline(y=0.1, line_dash="dash", line_color="orange", row=1, col=1)
fig.add_hline(y=0.2, line_dash="dash", line_color="red", row=1, col=1)
# P-value bar chart
colors_p = ['red' if p < 0.05 else 'green' for p in p_values]
fig.add_trace(
go.Bar(x=features, y=p_values, marker_color=colors_p, name='P-value'),
row=2, col=1
)
# Add significance threshold line
fig.add_hline(y=0.05, line_dash="dash", line_color="red", row=2, col=1)
fig.update_xaxes(tickangle=45)
fig.update_layout(height=800, title_text="Data Drift Analysis", showlegend=False)
fig.write_html(output_path)
print(f"[OK] Drift visualization saved to {output_path}")
# Example usage
if __name__ == "__main__":
# Load reference (training) data
reference_data = pd.read_csv('data/training_data.csv')
# Initialize detector
detector = DataDriftDetector(reference_data, threshold=0.05)
# Get current week's data
current_data = pd.read_csv('data/current_week_data.csv')
# Detect drift
results = detector.detect_drift(current_data)
# Print report
print(detector.get_drift_report(results))
# Generate visualization
detector.plot_drift(results, 'reports/drift_analysis.html')
# Alert if severe drift detected
severe_drift = [r for r in results if r.severity == 'severe']
if severe_drift:
print(f"\n[ERROR] SEVERE DRIFT DETECTED in {len(severe_drift)} features")
print("Consider retraining the model!")
# Send alert (integrate with your alerting system)
# send_alert(severity='critical', message=f'Severe data drift detected')Concept Drift Detection
Monitor changes in relationship between features and target variable.
Gama et al., 2014, ACM Computing Surveys - “A survey on concept drift adaptation”
from collections import deque
from sklearn.metrics import roc_auc_score, accuracy_score
import numpy as np
class ConceptDriftDetector:
"""
Detect concept drift by monitoring model performance over time
Uses ADWIN (Adaptive Windowing) algorithm to detect statistically significant
changes in error rate distribution
"""
def __init__(self, window_size: int = 100, min_window_size: int = 30):
"""
Args:
window_size: Maximum size of sliding window
min_window_size: Minimum samples needed for drift detection
"""
self.window_size = window_size
self.min_window_size = min_window_size
# Sliding window of recent performance
self.performance_window = deque(maxlen=window_size)
self.predictions_window = deque(maxlen=window_size)
self.actuals_window = deque(maxlen=window_size)
# Baseline performance
self.baseline_auc = None
self.drift_detected_count = 0
def update(self, y_true: np.ndarray, y_pred_proba: np.ndarray):
"""Add new batch of predictions with ground truth"""
# Calculate AUC for this batch
if len(np.unique(y_true)) > 1: # Need both classes
batch_auc = roc_auc_score(y_true, y_pred_proba)
self.performance_window.append(batch_auc)
# Store predictions and actuals for detailed analysis
self.predictions_window.extend(y_pred_proba)
self.actuals_window.extend(y_true)
def detect_drift(self, baseline_auc: float, threshold: float = 0.05) -> Dict:
"""
Detect if performance has degraded significantly
Args:
baseline_auc: Expected AUC from validation set
threshold: Acceptable drop in AUC (e.g., 0.05 = 5 percentage points)
Returns:
Dict with drift detection results
"""
if len(self.performance_window) < self.min_window_size:
return {
'drift_detected': False,
'message': 'Insufficient data for drift detection',
'samples_needed': self.min_window_size - len(self.performance_window)
}
# Calculate recent performance statistics
recent_aucs = list(self.performance_window)
mean_recent_auc = np.mean(recent_aucs)
std_recent_auc = np.std(recent_aucs)
min_recent_auc = np.min(recent_aucs)
# Performance drop
performance_drop = baseline_auc - mean_recent_auc
# Statistical test: Is recent performance significantly worse?
# Using one-sample t-test
from scipy import stats
t_statistic, p_value = stats.ttest_1samp(recent_aucs, baseline_auc)
# Drift detected if:
# 1. Mean performance dropped more than threshold
# 2. Difference is statistically significant
drift_detected = (performance_drop > threshold) and (p_value < 0.05) and (t_statistic < 0)
if drift_detected:
self.drift_detected_count += 1
# Calculate confidence interval
ci_95 = stats.t.interval(
0.95,
len(recent_aucs)-1,
loc=mean_recent_auc,
scale=stats.sem(recent_aucs)
)
return {
'drift_detected': drift_detected,
'baseline_auc': baseline_auc,
'current_auc': mean_recent_auc,
'current_auc_std': std_recent_auc,
'min_recent_auc': min_recent_auc,
'performance_drop': performance_drop,
'p_value': p_value,
'ci_95': ci_95,
'drift_count': self.drift_detected_count,
'message': self._generate_message(drift_detected, performance_drop, baseline_auc, mean_recent_auc)
}
def _generate_message(self, drift_detected: bool, drop: float, baseline: float, current: float) -> str:
"""Generate human-readable message"""
if not drift_detected:
return f"[OK] Performance stable (AUC: {current:.3f}, baseline: {baseline:.3f})"
else:
return f"[WARNING] CONCEPT DRIFT: Performance dropped from {baseline:.3f} to {current:.3f} (Δ={drop:.3f})"
def analyze_subgroup_drift(self, groups: pd.Series) -> Dict:
"""
Analyze if drift affects some subgroups more than others
Args:
groups: Protected attribute values for samples in window
"""
results = {}
y_true = np.array(list(self.actuals_window))
y_pred = np.array(list(self.predictions_window))
for group in groups.unique():
mask = (groups == group).values
if mask.sum() < 10: # Skip small groups
continue
group_auc = roc_auc_score(y_true[mask], y_pred[mask])
results[group] = {
'auc': group_auc,
'n_samples': mask.sum()
}
return results
# Example usage
detector = ConceptDriftDetector(window_size=100)
# Simulation: Collect predictions with ground truth over time
for week in range(1, 53): # 52 weeks
# Get predictions for this week
weekly_predictions = get_weekly_predictions(week)
# After 1 week, get ground truth (chart review confirms sepsis)
weekly_actuals = get_weekly_actuals(week)
# Update detector
detector.update(weekly_actuals, weekly_predictions)
# Check for drift every 4 weeks
if week % 4 == 0:
drift_status = detector.detect_drift(baseline_auc=0.85, threshold=0.05)
print(f"\nWeek {week} Drift Check:")
print(f" {drift_status['message']}")
if drift_status['drift_detected']:
print(f" [WARNING] Recommendation: Consider model retraining")
print(f" Current AUC: {drift_status['current_auc']:.3f} ± {drift_status['current_auc_std']:.3f}")
print(f" 95% CI: [{drift_status['ci_95'][0]:.3f}, {drift_status['ci_95'][1]:.3f}]")
# Alert ops team
send_alert(
severity='warning',
title='Concept Drift Detected',
message=drift_status['message']
)For comprehensive treatment of drift detection methods, see Lu et al., 2018, ACM Computing Surveys - “Learning under Concept Drift: A Review”.
Alerting System
Automated alerts when metrics exceed thresholds or anomalies detected.
import smtplib
from email.mime.text import MIMEText
from email.mime.multipart import MIMEMultipart
from typing import Dict, List, Optional
from datetime import datetime, timedelta
import requests
import json
class AlertManager:
"""
Comprehensive alerting system for ML model monitoring
Supports multiple channels: Email, Slack, PagerDuty
Implements alert cooldown to prevent spam
"""
def __init__(self, config: Dict):
self.config = config
self.alert_history = []
self.alert_cooldown = {} # Track last alert time per metric
def check_alerts(self, metrics: Dict):
"""
Check all alert conditions and send notifications
Args:
metrics: Dict containing current system/model metrics
"""
alerts = []
# Check prediction latency
if metrics.get('avg_latency', 0) > self.config['latency_threshold']:
alerts.append({
'severity': 'warning',
'title': 'High Prediction Latency',
'message': f"Average latency: {metrics['avg_latency']:.2f}s (threshold: {self.config['latency_threshold']}s)",
'metric': 'latency',
'details': {
'current': metrics['avg_latency'],
'threshold': self.config['latency_threshold'],
'p95': metrics.get('p95_latency'),
'p99': metrics.get('p99_latency')
}
})
# Check error rate
if metrics.get('error_rate', 0) > self.config['error_rate_threshold']:
alerts.append({
'severity': 'critical',
'title': 'High Error Rate',
'message': f"Error rate: {metrics['error_rate']:.2%} (threshold: {self.config['error_rate_threshold']:.2%})",
'metric': 'error_rate',
'details': {
'current': metrics['error_rate'],
'threshold': self.config['error_rate_threshold'],
'total_errors': metrics.get('total_errors'),
'total_requests': metrics.get('total_requests')
}
})
# Check data drift
if metrics.get('drift_detected', False):
drifted_features = metrics.get('drifted_features', [])
alerts.append({
'severity': 'warning',
'title': 'Data Drift Detected',
'message': f"Drift detected in {len(drifted_features)} features: {', '.join(drifted_features[:3])}{'...' if len(drifted_features) > 3 else ''}",
'metric': 'drift',
'details': {
'drifted_features': drifted_features,
'psi_scores': metrics.get('psi_scores', {})
}
})
# Check model performance degradation
if metrics.get('auc') and metrics['auc'] < self.config['min_auc']:
alerts.append({
'severity': 'critical',
'title': 'Model Performance Degradation',
'message': f"Current AUC: {metrics['auc']:.3f} (minimum: {self.config['min_auc']:.3f})",
'metric': 'performance',
'details': {
'current_auc': metrics['auc'],
'baseline_auc': self.config.get('baseline_auc'),
'min_auc': self.config['min_auc'],
'samples_evaluated': metrics.get('n_samples')
}
})
# Check prediction volume (unusually low might indicate system issue)
if metrics.get('predictions_per_minute', 0) < self.config.get('min_prediction_rate', 1):
alerts.append({
'severity': 'warning',
'title': 'Low Prediction Volume',
'message': f"Only {metrics['predictions_per_minute']:.1f} predictions/min",
'metric': 'volume',
'details': {
'current_rate': metrics['predictions_per_minute'],
'expected_min': self.config.get('min_prediction_rate')
}
})
# Send all triggered alerts
for alert in alerts:
self._send_alert(alert)
def _send_alert(self, alert: Dict):
"""
Send alert via configured channels with cooldown logic
Args:
alert: Dict with alert details
"""
metric = alert['metric']
# Check cooldown to prevent spam
if metric in self.alert_cooldown:
last_alert = self.alert_cooldown[metric]
cooldown_seconds = self.config.get('alert_cooldown_seconds', 3600)
if (datetime.now() - last_alert).seconds < cooldown_seconds:
return # Skip this alert, still in cooldown
# Send via configured channels
if self.config.get('email_alerts', False):
self._send_email(alert)
if self.config.get('slack_alerts', False):
self._send_slack(alert)
if self.config.get('pagerduty_alerts', False) and alert['severity'] == 'critical':
self._send_pagerduty(alert)
# Log alert
alert_record = {
**alert,
'timestamp': datetime.now(),
'channels_notified': []
}
if self.config.get('email_alerts'):
alert_record['channels_notified'].append('email')
if self.config.get('slack_alerts'):
alert_record['channels_notified'].append('slack')
if self.config.get('pagerduty_alerts') and alert['severity'] == 'critical':
alert_record['channels_notified'].append('pagerduty')
self.alert_history.append(alert_record)
# Update cooldown
self.alert_cooldown[metric] = datetime.now()
def _send_email(self, alert: Dict):
"""Send alert via email"""
try:
msg = MIMEMultipart('alternative')
msg['Subject'] = f"[{alert['severity'].upper()}] {alert['title']}"
msg['From'] = self.config['email_from']
msg['To'] = ', '.join(self.config['email_recipients'])
# Plain text version
text = f"""
{alert['title']}
Severity: {alert['severity'].upper()}
{alert['message']}
Details:
{json.dumps(alert.get('details', {}), indent=2)}
Timestamp: {datetime.now().isoformat()}
--
Automated alert from ML Monitoring System
"""
# HTML version
severity_color = '#dc3545' if alert['severity'] == 'critical' else '#ffc107'
html = f"""
<html>
<head></head>
<body>
<div style="font-family: Arial, sans-serif; padding: 20px;">
<h2 style="color: {severity_color};">{alert['title']}</h2>
<p><strong>Severity:</strong> <span style="color: {severity_color}; font-weight: bold;">{alert['severity'].upper()}</span></p>
<p>{alert['message']}</p>
<h3>Details:</h3>
<pre style="background-color: #f5f5f5; padding: 10px; border-radius: 5px;">{json.dumps(alert.get('details', {}), indent=2)}</pre>
<p style="color: #666; font-size: 12px;">Timestamp: {datetime.now().isoformat()}</p>
<hr>
<p style="color: #999; font-size: 11px;">Automated alert from ML Monitoring System</p>
</div>
</body>
</html>
"""
part1 = MIMEText(text, 'plain')
part2 = MIMEText(html, 'html')
msg.attach(part1)
msg.attach(part2)
# Send email
with smtplib.SMTP(self.config['smtp_server'], self.config.get('smtp_port', 587)) as server:
if self.config.get('smtp_use_tls', True):
server.starttls()
if self.config.get('smtp_username') and self.config.get('smtp_password'):
server.login(self.config['smtp_username'], self.config['smtp_password'])
server.send_message(msg)
print(f"[OK] Email alert sent: {alert['title']}")
except Exception as e:
print(f"[ERROR] Failed to send email alert: {e}")
def _send_slack(self, alert: Dict):
"""Send alert via Slack webhook"""
try:
color = 'danger' if alert['severity'] == 'critical' else 'warning'
payload = {
'attachments': [{
'color': color,
'title': alert['title'],
'text': alert['message'],
'fields': [
{
'title': key.replace('_', ' ').title(),
'value': str(value),
'short': True
}
for key, value in alert.get('details', {}).items()
],
'footer': 'ML Model Monitoring',
'ts': int(datetime.now().timestamp())
}]
}
response = requests.post(
self.config['slack_webhook_url'],
json=payload,
timeout=10
)
if response.status_code == 200:
print(f"[OK] Slack alert sent: {alert['title']}")
else:
print(f"[ERROR] Slack alert failed: {response.status_code}")
except Exception as e:
print(f"[ERROR] Failed to send Slack alert: {e}")
def _send_pagerduty(self, alert: Dict):
"""Send critical alert to PagerDuty"""
try:
payload = {
'routing_key': self.config['pagerduty_integration_key'],
'event_action': 'trigger',
'payload': {
'summary': alert['title'],
'severity': alert['severity'],
'source': 'ML Monitoring System',
'custom_details': alert.get('details', {})
}
}
response = requests.post(
'https://events.pagerduty.com/v2/enqueue',
json=payload,
timeout=10
)
if response.status_code == 202:
print(f"[OK] PagerDuty alert sent: {alert['title']}")
else:
print(f"[ERROR] PagerDuty alert failed: {response.status_code}")
except Exception as e:
print(f"[ERROR] Failed to send PagerDuty alert: {e}")
def get_alert_summary(self, hours: int = 24) -> Dict:
"""Get summary of alerts in past N hours"""
cutoff = datetime.now() - timedelta(hours=hours)
recent_alerts = [a for a in self.alert_history if a['timestamp'] > cutoff]
summary = {
'total_alerts': len(recent_alerts),
'critical': sum(1 for a in recent_alerts if a['severity'] == 'critical'),
'warning': sum(1 for a in recent_alerts if a['severity'] == 'warning'),
'by_metric': {}
}
for alert in recent_alerts:
metric = alert['metric']
summary['by_metric'][metric] = summary['by_metric'].get(metric, 0) + 1
return summary
# Example configuration
alert_config = {
# Thresholds
'latency_threshold': 1.0, # seconds
'error_rate_threshold': 0.05, # 5%
'min_auc': 0.75,
'baseline_auc': 0.85,
'min_prediction_rate': 1.0, # predictions per minute
# Cooldown
'alert_cooldown_seconds': 3600, # 1 hour between same alert
# Email configuration
'email_alerts': True,
'email_from': 'mlops@hospital.org',
'email_recipients': ['datascience-team@hospital.org', 'clinical-ops@hospital.org'],
'smtp_server': 'smtp.hospital.org',
'smtp_port': 587,
'smtp_use_tls': True,
'smtp_username': 'mlops@hospital.org',
'smtp_password': 'secure_password', # Use secrets management in production
# Slack configuration
'slack_alerts': True,
'slack_webhook_url': 'https://hooks.slack.com/services/YOUR/WEBHOOK/URL',
# PagerDuty configuration (for critical alerts)
'pagerduty_alerts': True,
'pagerduty_integration_key': 'your_integration_key'
}
alert_manager = AlertManager(alert_config)
# Check alerts periodically (e.g., every 5 minutes via cron/Airflow)
def monitor_and_alert():
"""Collect metrics and check for alert conditions"""
# Collect current metrics
metrics = {
'avg_latency': get_avg_latency_last_5min(),
'p95_latency': get_p95_latency_last_5min(),
'p99_latency': get_p99_latency_last_5min(),
'error_rate': get_error_rate_last_5min(),
'total_errors': get_total_errors_last_5min(),
'total_requests': get_total_requests_last_5min(),
'predictions_per_minute': get_prediction_rate(),
'drift_detected': check_data_drift(),
'drifted_features': get_drifted_features() if check_data_drift() else [],
'psi_scores': get_psi_scores(),
'auc': get_current_auc(),
'n_samples': get_evaluated_sample_count()
}
# Check and send alerts
alert_manager.check_alerts(metrics)
# Schedule: */5 * * * * (every 5 minutes)