from sklearn.model_selection import cross_val_score from sklearn.metrics import roc_auc_score class AdversarialDriftDetector(BaseEstimator, TransformerMixin): """ Trains a classifier to tell reference data apart from new data. If it can, something about the joint distribution moved - not necessarily any single feature. """ def __init__(self, auc_threshold=0.65, cv_folds=5, min_rows_for_cv=200, random_state=42): self.auc_threshold = auc_threshold self.cv_folds = cv_folds # below this, cross_val_score is more likely to error on a thin # fold than tell you anything useful - learned that against a # 90-row canary batch that kept failing for no obvious reason self.min_rows_for_cv = min_rows_for_cv self.random_state = random_state def fit(self, X, y=None): self.ref_df_ = X.select_dtypes(include=np.number).copy() self.ref_cols_ = self.ref_df_.columns.tolist() return self def check(self, X_new): missing = set(self.ref_cols_) - set(X_new.columns) if missing: raise ValueError(f"X_new is missing columns the detector was fit on: {missing}") new_df = X_new[self.ref_cols_].copy() if new_df.isna().any().any(): # don't silently drop rows here - a batch that's suddenly full # of nulls is itself a drift signal worth knowing about logger.warning( "check() got %d rows with NaNs, filling with column medians", new_df.isna().any(axis=1).sum(), ) new_df = new_df.fillna(self.ref_df_.median()) combined = pd.concat([self.ref_df_, new_df], ignore_index=True) labels = np.r_[np.zeros(len(self.ref_df_)), np.ones(len(new_df))] rf = RandomForestClassifier(n_estimators=100, n_jobs=-1, random_state=self.random_state) if len(new_df) < self.min_rows_for_cv: Xtr, Xte, ytr, yte = train_test_split( combined, labels, test_size=0.3, random_state=self.random_state, stratify=labels ) rf.fit(Xtr, ytr) auc = roc_auc_score(yte, rf.predict_proba(Xte)[:, 1]) else: auc = cross_val_score(rf, combined, labels, cv=self.cv_folds, scoring="roc_auc").mean() drifted = auc > self.auc_threshold top_features = None if drifted: rf.fit(combined, labels) top_features = pd.Series(rf.feature_importances_, index=self.ref_cols_) \ .sort_values(ascending=False).head(3) logger.warning("Drift detected, auc=%.3f, top features: %s", auc, top_features.index.tolist()) return {"auc": auc, "drift_detected": drifted, "top_drift_features": top_features} def transform(self, X): return X