Once a machine learning model is deployed to production, input features tend to shift away from the training distribution due to seasonal trends, user behavior changes, or upstream pipeline alterations. Identifying this data drift before model accuracy degrades is crucial for maintaining AI service stability.
How to detect data drift in production Machine Learning models using Python?
1 Answer
Data drift is a change in the inputs a model receives after deployment. A model trained on last year’s customer mix, for example, may encounter a very different mix this month. Comparing recent production data with a representative reference dataset can help you spot that change before it quietly affects predictions.
Here’s a small batch check for numerical and categorical features using pandas and SciPy. It compares non-missing numerical values with the Kolmogorov–Smirnov test and categorical distributions with a chi-square test. It also reports changes in missing-value rates.
import pandas as pd
from scipy.stats import chi2_contingency, ks_2samp
def detect_drift(reference, current, numeric_features, categorical_features,
p_threshold=0.01):
# Keep the reference period representative of normal production inputs.
results = []
for feature in numeric_features:
ref = reference[feature].dropna()
cur = current[feature].dropna()
# Skip tests that cannot be interpreted with too few observations.
if len(ref) < 2 or len(cur) < 2:
continue
test = ks_2samp(ref, cur)
ref_missing = reference[feature].isna().mean()
cur_missing = current[feature].isna().mean()
results.append({
"feature": feature,
"test": "KS",
"statistic": test.statistic,
"p_value": test.pvalue,
"missing_rate_change": cur_missing - ref_missing,
"drift_flag": test.pvalue < p_threshold,
})
for feature in categorical_features:
# Treat missing values as a category so their frequency is compared too.
ref = reference[feature].fillna("__MISSING__").astype(str)
cur = current[feature].fillna("__MISSING__").astype(str)
categories = sorted(set(ref) | set(cur))
# Align category counts before testing the two distributions.
counts = pd.DataFrame({
"reference": ref.value_counts().reindex(categories, fill_value=0),
"current": cur.value_counts().reindex(categories, fill_value=0),
})
if len(categories) < 2:
continue
test = chi2_contingency(counts.T)
results.append({
"feature": feature,
"test": "chi-square",
"statistic": test.statistic,
"p_value": test.pvalue,
"drift_flag": test.pvalue < p_threshold,
})
return pd.DataFrame(results)
# Example: compare a recent production batch with a saved reference sample.
report = detect_drift(
reference=reference_df,
current=production_batch,
numeric_features=["age", "account_balance"],
categorical_features=["country", "plan_type"],
)
print(report.sort_values("p_value"))
Save a reference sample from a period when the model was behaving acceptably, then run the comparison on consistent production windows—daily or weekly, depending on traffic. The p-value threshold above is only an example. With many features or very large batches, tiny changes can become statistically significant, so review the test statistic, sample size, missingness, and operational impact before paging someone.
This checks input distributions, not whether predictions are still correct. When labels become available, monitor model quality separately against those outcomes. A drift alert is a prompt to investigate; it is not, by itself, a reason to retrain.