PIVOT / src /pivot /evaluation /metrics.py
pranamanam's picture
Upload 176 files
6fa9282 verified
Raw History Blame Contribute Delete
2.45 kB
"""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),
}