File size: 2,446 Bytes
6fa9282
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
"""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),
    }