EyeQC / src /batch_effects.py
rdharini2001's picture
EyeQC
6fe482b
Raw
History Blame Contribute Delete
6.87 kB
"""
Batch-effect analysis for retinal image cohorts.
When fundus images come from different cameras, sites or acquisition days, they
carry systematic technical variation ("batch effects") that can dominate the
biological signal - a well-known confounder in imaging cohorts and foundation-
model embeddings. This module provides:
* feature extraction (FLAIR embeddings when available, else interpretable
QC + colour descriptors),
* batch-effect *detection* (a batch-classifier AUC + silhouette by batch),
* batch-effect *visualisation* (PCA / optional UMAP scatter, per-feature
distributions),
* batch-effect *correction* via a self-contained ComBat (empirical-Bayes
location/scale harmonisation) plus a simple z-standardisation option.
"""
from __future__ import annotations
import numpy as np
from sklearn.decomposition import PCA
from sklearn.preprocessing import StandardScaler
from sklearn.linear_model import LogisticRegression
from sklearn.model_selection import cross_val_predict
from sklearn.metrics import roc_auc_score, silhouette_score
# ----------------------------------------------------------------- feature extraction
def descriptor_from_metrics(rgb, metrics, fov):
"""Interpretable per-image feature vector: the QC scores plus low-order colour
/ intensity statistics. Used as the embedding when FLAIR is not loaded."""
import cv2
from .qc_metrics import inner_disc_mask
m = inner_disc_mask(rgb.shape, fov, 0.9)
if m.sum() < 50:
m = np.ones(rgb.shape[:2], bool)
feats = [mm["score"] for mm in metrics]
for c in range(3):
ch = rgb[..., c][m].astype(np.float32)
feats += [ch.mean() / 255, ch.std() / 255]
hsv = cv2.cvtColor(rgb, cv2.COLOR_RGB2HSV)
feats += [hsv[..., 0][m].mean() / 180, hsv[..., 1][m].mean() / 255]
feats += [fov["coverage"], fov["radius"] / max(rgb.shape[:2])]
return np.array(feats, np.float32)
# ------------------------------------------------------------------------ detection
def detect_batch_effect(X, batches):
"""Quantify how strongly `batches` are separable in feature space.
Returns dict with:
auc : cross-validated one-vs-rest batch-classifier AUC (0.5 == none)
silhouette : silhouette of the batch labelling in PCA space (-1..1)
severity : qualitative label
"""
X = np.asarray(X, np.float32)
b = np.asarray(batches)
uniq = np.unique(b)
out = dict(auc=np.nan, silhouette=np.nan, severity="n/a")
if len(uniq) < 2 or X.shape[0] < 6:
out["severity"] = "insufficient data (need >=2 batches, >=6 images)"
return out
Xs = StandardScaler().fit_transform(X)
# Batch classifier AUC (macro one-vs-rest)
try:
y = np.searchsorted(uniq, b)
clf = LogisticRegression(max_iter=1000)
proba = cross_val_predict(clf, Xs, y, cv=min(4, len(uniq) + 1),
method="predict_proba")
if proba.shape[1] == 2:
auc = roc_auc_score(y, proba[:, 1])
else:
auc = roc_auc_score(y, proba, multi_class="ovr", average="macro")
out["auc"] = float(auc)
except Exception:
pass
try:
emb = PCA(n_components=min(10, Xs.shape[1])).fit_transform(Xs)
out["silhouette"] = float(silhouette_score(emb, b))
except Exception:
pass
auc = out["auc"]
if not np.isnan(auc):
out["severity"] = ("negligible" if auc < 0.6 else
"mild" if auc < 0.75 else
"moderate" if auc < 0.9 else "severe")
return out
def embed_2d(X, method="pca", seed=0):
"""2-D embedding for scatter plots. UMAP if available, else PCA."""
X = StandardScaler().fit_transform(np.asarray(X, np.float32))
if method == "umap":
try:
import umap
n = X.shape[0]
reducer = umap.UMAP(n_neighbors=min(15, max(2, n - 1)),
min_dist=0.1, random_state=seed)
return reducer.fit_transform(X), "UMAP"
except Exception:
pass
return PCA(n_components=2).fit_transform(X), "PCA"
# ------------------------------------------------------------------------ correction
def combat(X, batches, covariates=None):
"""Self-contained parametric ComBat harmonisation.
Removes additive (location) and multiplicative (scale) batch effects with an
empirical-Bayes shrinkage, optionally preserving biological covariates.
Follows Johnson et al. 2007. X is (n_samples, n_features).
"""
X = np.asarray(X, np.float64)
b = np.asarray(batches)
n, p = X.shape
uniq = np.unique(b)
if len(uniq) < 2:
return X.astype(np.float32)
# Design: intercept + optional covariates
if covariates is not None and len(np.unique(covariates)) > 1:
cov = np.asarray(covariates)
cats = np.unique(cov)
C = np.column_stack([(cov == c).astype(float) for c in cats])
design = C # already spans intercept
else:
design = np.ones((n, 1))
# Standardise using pooled (grand) model
B = np.linalg.lstsq(design, X, rcond=None)[0]
grand = design @ B
var = ((X - grand) ** 2).mean(axis=0) + 1e-8
sd = np.sqrt(var)
Z = (X - grand) / sd
Xc = X.copy()
for g in uniq:
idx = np.where(b == g)[0]
if len(idx) < 2:
continue
Zi = Z[idx]
gamma_hat = Zi.mean(axis=0) # additive batch shift
delta_hat = Zi.var(axis=0) + 1e-8 # multiplicative batch scale
# Empirical-Bayes priors across features
gbar, t2 = gamma_hat.mean(), gamma_hat.var() + 1e-8
dbar, dvar = delta_hat.mean(), delta_hat.var() + 1e-8
a = (2 * dvar + dbar ** 2) / dvar
bpr = (dbar * dvar + dbar ** 3) / dvar
# Iterative EB estimates (few fixed-point iterations)
gamma_star, delta_star = gamma_hat.copy(), delta_hat.copy()
for _ in range(30):
gamma_star = (len(idx) * t2 * gamma_hat + delta_star * gbar) / \
(len(idx) * t2 + delta_star)
resid = ((Zi - gamma_star) ** 2).sum(axis=0)
delta_star = (0.5 * resid + bpr) / (0.5 * len(idx) + a - 1)
delta_star = np.clip(delta_star, 1e-6, None)
Zi_adj = (Zi - gamma_star) / np.sqrt(delta_star)
Xc[idx] = Zi_adj * sd + grand[idx]
return Xc.astype(np.float32)
def zstandardise_by_batch(X, batches):
"""Simpler baseline: per-batch z-standardisation to a common mean/var."""
X = np.asarray(X, np.float64).copy()
b = np.asarray(batches)
gmean, gsd = X.mean(0), X.std(0) + 1e-8
for g in np.unique(b):
idx = b == g
mu, sd = X[idx].mean(0), X[idx].std(0) + 1e-8
X[idx] = (X[idx] - mu) / sd * gsd + gmean
return X.astype(np.float32)