"""Out-of-distribution scoring and abstention.
Convention: every scorer returns a 1-D array where HIGHER = more OOD.
"""
from __future__ import annotations
from abc import ABC, abstractmethod
import numpy as np
from scipy.special import logsumexp
from sklearn.covariance import LedoitWolf
from sklearn.decomposition import PCA
from sklearn.neighbors import NearestNeighbors
from sklearn.pipeline import Pipeline
from sklearn.preprocessing import StandardScaler
[docs]
class OODScorer(ABC):
name: str = "ood"
[docs]
@abstractmethod
def fit(self, X: np.ndarray, y: np.ndarray | None = None, model=None) -> "OODScorer":
...
[docs]
@abstractmethod
def score(self, X: np.ndarray, model=None) -> np.ndarray:
...
[docs]
class MaxSoftmax(OODScorer):
"""1 - max predicted probability."""
name = "max_softmax"
[docs]
def fit(self, X, y=None, model=None):
return self
[docs]
def score(self, X, model=None):
return 1.0 - model.predict_proba(X).max(axis=1)
[docs]
class Entropy(OODScorer):
"""Normalized Shannon entropy of the predictive distribution."""
name = "entropy"
[docs]
def fit(self, X, y=None, model=None):
return self
[docs]
def score(self, X, model=None):
proba = np.clip(model.predict_proba(X), 1e-12, 1.0)
return -np.sum(proba * np.log(proba), axis=1) / np.log(proba.shape[1])
[docs]
class Energy(OODScorer):
"""Negative log-sum-exp of logits (log-proba fallback for CV-calibrated models)."""
name = "energy"
[docs]
def fit(self, X, y=None, model=None):
return self
[docs]
def score(self, X, model=None):
if hasattr(model, "_logits"):
logits = model._logits(X)
else:
logits = np.log(np.clip(model.predict_proba(X), 1e-12, 1.0))
return -logsumexp(logits, axis=1)
class _PCASpace:
"""Shared scaler+PCA embedding for feature-space detectors."""
def __init__(self, n_components: int, random_state: int = 42) -> None:
self.pipeline = Pipeline([
("scaler", StandardScaler()),
("pca", PCA(n_components=n_components, random_state=random_state)),
])
def fit_transform(self, X: np.ndarray) -> np.ndarray:
return self.pipeline.fit_transform(X)
def transform(self, X: np.ndarray) -> np.ndarray:
return self.pipeline.transform(X)
[docs]
class Mahalanobis(OODScorer):
"""Min per-class Mahalanobis distance in PCA space (shared LW covariance)."""
name = "mahalanobis"
[docs]
def __init__(self, n_components: int = 20, random_state: int = 42) -> None:
self.space = _PCASpace(n_components, random_state)
[docs]
def fit(self, X, y=None, model=None):
Z = self.space.fit_transform(X)
y = np.asarray(y)
self.means_ = np.stack([Z[y == c].mean(axis=0) for c in np.unique(y)])
pooled = np.concatenate([Z[y == c] - Z[y == c].mean(axis=0) for c in np.unique(y)])
self.precision_ = LedoitWolf().fit(pooled).get_precision()
return self
[docs]
def score(self, X, model=None):
Z = self.space.transform(X)
dists = []
for mean in self.means_:
diff = Z - mean
dists.append(np.einsum("ij,jk,ik->i", diff, self.precision_, diff))
return np.sqrt(np.clip(np.min(np.stack(dists), axis=0), 0.0, None))
[docs]
class KNNDistance(OODScorer):
"""Distance to the k-th nearest training sample in PCA space."""
name = "knn"
[docs]
def __init__(self, k: int = 5, n_components: int = 20, random_state: int = 42) -> None:
self.k = k
self.space = _PCASpace(n_components, random_state)
[docs]
def fit(self, X, y=None, model=None):
Z = self.space.fit_transform(X)
self.nn_ = NearestNeighbors(n_neighbors=self.k).fit(Z)
return self
[docs]
def score(self, X, model=None):
dist, _ = self.nn_.kneighbors(self.space.transform(X))
return dist[:, -1]
[docs]
class SpectralAngleScorer(OODScorer):
"""Min spectral angle to the class-mean spectra (norm route only)."""
name = "spectral_angle"
[docs]
def fit(self, X, y=None, model=None):
y = np.asarray(y)
self.means_ = np.stack([X[y == c].mean(axis=0) for c in np.unique(y)])
return self
[docs]
def score(self, X, model=None):
from xpectrass import spectral_angle
angles = np.stack([
np.array([spectral_angle(x, mean) for x in X]) for mean in self.means_
])
return angles.min(axis=0)
ALL_SCORERS = (MaxSoftmax, Entropy, Energy, Mahalanobis, KNNDistance, SpectralAngleScorer)
[docs]
class AbstainingClassifier:
"""Predict a class when the OOD score is below tau, otherwise abstain."""
ABSTAIN = "ABSTAIN"
[docs]
def __init__(self, model, scorer: OODScorer) -> None:
self.model = model
self.scorer = scorer
self.tau_: float | None = None
[docs]
def fit_tau(self, X_val: np.ndarray, target_coverage: float = 0.90) -> "AbstainingClassifier":
"""Set tau so that `target_coverage` of X_val falls below it.
X_val must be a HELD-OUT labeled set that neither the model nor the
scorer was trained on — calibrating on in-sample data yields optimistic
scores and a threshold too tight to hit the target coverage on new data.
"""
scores = self.scorer.score(X_val, model=self.model)
self.tau_ = float(np.quantile(scores, target_coverage))
return self
[docs]
def achieved_coverage(self, X: np.ndarray) -> float:
"""Fraction of X that would be classified (not abstained) at the current tau."""
if self.tau_ is None:
raise RuntimeError("Call fit_tau before achieved_coverage.")
return float(np.mean(self.scorer.score(X, model=self.model) < self.tau_))
[docs]
def predict(self, X: np.ndarray, class_names: list[str]) -> tuple[np.ndarray, np.ndarray]:
if self.tau_ is None:
raise RuntimeError("Call fit_tau before predict.")
scores = self.scorer.score(X, model=self.model)
labels = np.array([class_names[i] for i in self.model.predict(X)], dtype=object)
labels[scores >= self.tau_] = self.ABSTAIN
return labels, scores
# --- Synthetic OOD perturbations -------------------------------------------
[docs]
def add_noise(X: np.ndarray, sigma: float, rng: np.random.Generator) -> np.ndarray:
scale = np.std(X, axis=1, keepdims=True)
return X + rng.normal(0.0, sigma, X.shape) * scale
[docs]
def baseline_drift(X: np.ndarray, magnitude: float, rng: np.random.Generator) -> np.ndarray:
n, p = X.shape
ramp = np.linspace(0.0, 1.0, p)[None, :]
slope = rng.uniform(-magnitude, magnitude, (n, 1)) * np.std(X, axis=1, keepdims=True)
return X + slope * ramp
[docs]
def block_scramble(X: np.ndarray, n_blocks: int, rng: np.random.Generator) -> np.ndarray:
"""Shuffle contiguous wavenumber blocks: destroys band structure, keeps
the overall intensity distribution."""
p = X.shape[1]
edges = np.linspace(0, p, n_blocks + 1, dtype=int)
blocks = [X[:, edges[i]:edges[i + 1]] for i in range(n_blocks)]
order = rng.permutation(n_blocks)
return np.concatenate([blocks[i] for i in order], axis=1)