Distributed E-Value Prediction Sets / distributed_evalue_experiment.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  1import json
  2from pathlib import Path
  3import numpy as np
  4
  5SEED = 939
  6K, M = 5, 5
  7N_CAL, N_TEST = 12000, 50000
  8ALPHA = 0.10
  9
 10
 11def softmax(z):
 12    z = z - z.max(axis=1, keepdims=True)
 13    q = np.exp(z)
 14    return q / q.sum(axis=1, keepdims=True)
 15
 16
 17def make_probs(y, signal, noise, rng):
 18    logits = rng.normal(0.0, noise, size=(len(y), K))
 19    logits[np.arange(len(y)), y] += signal
 20    return softmax(logits)
 21
 22
 23def e_values(probs, normalizer, beta=1.0):
 24    # Evidence against candidate y: low model probability => high e-value.
 25    # Calibration normalizer enforces empirical E[e_i(X,Y)] ~= 1.
 26    return (1.0 / (K * probs)) ** beta / normalizer
 27
 28
 29def summarize(e, y, weights, alpha):
 30    fused = np.sum(e * weights[None, None, :], axis=2)
 31    sets = fused < 1.0 / alpha
 32    true_e = fused[np.arange(len(y)), y]
 33    return {
 34        "mean_fused_true_e": float(true_e.mean()),
 35        "coverage": float(sets[np.arange(len(y)), y].mean()),
 36        "rejection": float((true_e >= 1.0 / alpha).mean()),
 37        "markov_bound": float(alpha * true_e.mean()),
 38        "avg_set_size": float(sets.sum(axis=1).mean()),
 39        "singleton_rate": float((sets.sum(axis=1) == 1).mean()),
 40        "point_accuracy": float((np.argmin(fused, axis=1) == y).mean()),
 41    }
 42
 43
 44def main():
 45    rng = np.random.default_rng(SEED)
 46    signals = np.array([1.25, 1.55, 1.85, 2.15, 2.45])
 47    noises = np.array([1.25, 1.10, 0.95, 0.85, 0.75])
 48    uncertainty = np.array([1.00, 0.72, 0.48, 0.28, 0.12])
 49    ycal, ytest = rng.integers(K, size=N_CAL), rng.integers(K, size=N_TEST)
 50    cal_probs, test_probs, normalizers = [], [], []
 51    for signal, noise in zip(signals, noises):
 52        pc = make_probs(ycal, signal, noise, rng)
 53        pt = make_probs(ytest, signal, noise, rng)
 54        raw = e_values(pc, 1.0)[np.arange(N_CAL), ycal]
 55        normalizers.append(raw.mean() * 1.01)  # finite-sample safety margin
 56        cal_probs.append(pc); test_probs.append(pt)
 57    normalizers = np.asarray(normalizers)
 58    test_e = np.stack([e_values(p, z) for p, z in zip(test_probs, normalizers)], axis=2)
 59    cal_e_true = np.stack([e_values(p, z)[np.arange(N_CAL), ycal]
 60                           for p, z in zip(cal_probs, normalizers)], axis=1)
 61    local_cal = cal_e_true.mean(axis=0)
 62    local_test = np.array([test_e[np.arange(N_TEST), ytest, i].mean() for i in range(M)])
 63    out = {"seed": SEED, "n_cal": N_CAL, "n_test": N_TEST, "alpha": ALPHA,
 64           "construction": "inverse_probability_e_value", "uncertainty": uncertainty.tolist(),
 65           "normalizers": normalizers.tolist(), "local_calibration_true_e_means": local_cal.tolist(),
 66           "local_test_true_e_means": local_test.tolist(), "sweeps": {}}
 67
 68    # Prediction 1: E[e_F] is the convex weighted mean, for all kappa.
 69    out["sweeps"]["kappa"] = []
 70    for kappa in [0.0, 0.5, 1.0, 2.0, 5.0]:
 71        w = np.exp(-kappa * uncertainty); w /= w.sum()
 72        row = summarize(test_e, ytest, w, ALPHA)
 73        row.update(kappa=kappa, weights=w.tolist(), predicted_mean=float(w @ local_test))
 74        row["convexity_abs_error"] = abs(row["mean_fused_true_e"] - row["predicted_mean"])
 75        out["sweeps"]["kappa"].append(row)
 76
 77    # Prediction 2: every active neighborhood remains valid; size can improve.
 78    out["sweeps"]["active_experts"] = []
 79    for n in range(1, M + 1):
 80        w = np.exp(-2.0 * uncertainty[:n]); w /= w.sum()
 81        row = summarize(test_e[:, :, :n], ytest, w, ALPHA)
 82        row.update(n_active=n, weights=w.tolist(), predicted_mean=float(w @ local_test[:n]))
 83        out["sweeps"]["active_experts"].append(row)
 84
 85    # Prediction 3: Markov rejection bound is alpha * E[e_F].
 86    out["sweeps"]["alpha"] = []
 87    w = np.exp(-2.0 * uncertainty); w /= w.sum()
 88    for alpha in [0.02, 0.05, 0.10, 0.20, 0.40]:
 89        row = summarize(test_e, ytest, w, alpha); row.update(alpha=alpha, weights=w.tolist())
 90        out["sweeps"]["alpha"].append(row)
 91
 92    best_p = test_probs[-1]
 93    best_e = e_values(best_p, normalizers[-1])
 94    base_sets = best_e < 1.0 / ALPHA
 95    out["baseline"] = {
 96        "strongest_expert_point_accuracy": float((np.argmax(best_p, axis=1) == ytest).mean()),
 97        "strongest_expert_coverage": float(base_sets[np.arange(N_TEST), ytest].mean()),
 98        "strongest_expert_avg_set_size": float(base_sets.sum(axis=1).mean()),
 99        "unweighted": summarize(test_e, ytest, np.ones(M) / M, ALPHA),
100        "uncertainty_weighted": summarize(test_e, ytest, w, ALPHA),
101    }
102    Path("results.json").write_text(json.dumps(out, indent=2))
103    print(json.dumps(out, indent=2))
104
105
106if __name__ == "__main__": main()