"""Population and retrieval metrics with explicit feature and target units.""" from __future__ import annotations import numpy as np from scipy.spatial.distance import cdist from scipy.stats import pearsonr, wasserstein_distance def correlation(pred, observed) -> float: """Pearson correlation across features; undefined constant vectors return NaN.""" p = np.asarray(pred).ravel() y = np.asarray(observed).ravel() if p.std() < 1e-12 or y.std() < 1e-12: return float("nan") return float(pearsonr(p, y).statistic) def mmd2(x, y, gamma: float) -> float: """Biased RBF MMD squared. gamma must be fixed across compared methods.""" if gamma <= 0: raise ValueError("gamma must be positive") k = lambda a, b: np.exp(-gamma * cdist(a, b, metric="sqeuclidean")) return float(k(x, x).mean() + k(y, y).mean() - 2 * k(x, y).mean()) def sliced_wasserstein(x, y, n_proj=50, seed=0) -> float: """Mean exact 1-D Wasserstein-1 distance over reproducible unit directions.""" rng = np.random.default_rng(seed) w = rng.normal(size=(n_proj, x.shape[1])) w /= np.linalg.norm(w, axis=1, keepdims=True) return float(np.mean([wasserstein_distance(x @ v, y @ v) for v in w])) def energy_distance(x, y) -> float: """Empirical energy statistic in PCA coordinates, including diagonal terms.""" return float(2 * cdist(x, y).mean() - cdist(x, x).mean() - cdist(y, y).mean()) def retrieval_metrics(ranked, truth, k=5) -> dict: """Exact recovery and one-relevant-item nDCG; censored ranks stay undefined.""" rank = list(ranked).index(truth) + 1 if truth in ranked else None return { "top1": float(rank == 1), "top5": float(rank is not None and rank <= k), "ndcg10": ( float(1 / np.log2(rank + 1)) if rank is not None and rank <= 10 else 0.0 ), "rank": rank, } def bootstrap(values, seed=0, n_boot=2000) -> dict: """Percentile interval over independent conditions, never over cells.""" x = np.asarray(values, float) x = x[np.isfinite(x)] if not len(x): return {"mean": None, "lower": None, "upper": None, "n": 0} rng = np.random.default_rng(seed) means = x[rng.integers(len(x), size=(n_boot, len(x)))].mean(1) lo, hi = np.quantile(means, [0.025, 0.975]) return { "mean": float(x.mean()), "lower": float(lo), "upper": float(hi), "n": len(x), }