Download src/pivot/evaluation/metrics.py from ChatterjeeLab/PIVOT: direct link, hf CLI and curl.
- Browser
- Download file 2.45 kB
-
https://huggingface.co/ChatterjeeLab/PIVOT/resolve/main/src/pivot/evaluation/metrics.py
- Command line
-
hf download hf://ChatterjeeLab/PIVOT/src/pivot/evaluation/metrics.py
-
curl -L -o metrics.py https://huggingface.co/ChatterjeeLab/PIVOT/resolve/main/src/pivot/evaluation/metrics.py
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), | |
| } | |