import json from pathlib import Path import numpy as np SEED = 939 K, M = 5, 5 N_CAL, N_TEST = 12000, 50000 ALPHA = 0.10 def softmax(z): z = z - z.max(axis=1, keepdims=True) q = np.exp(z) return q / q.sum(axis=1, keepdims=True) def make_probs(y, signal, noise, rng): logits = rng.normal(0.0, noise, size=(len(y), K)) logits[np.arange(len(y)), y] += signal return softmax(logits) def e_values(probs, normalizer, beta=1.0): # Evidence against candidate y: low model probability => high e-value. # Calibration normalizer enforces empirical E[e_i(X,Y)] ~= 1. return (1.0 / (K * probs)) ** beta / normalizer def summarize(e, y, weights, alpha): fused = np.sum(e * weights[None, None, :], axis=2) sets = fused < 1.0 / alpha true_e = fused[np.arange(len(y)), y] return { "mean_fused_true_e": float(true_e.mean()), "coverage": float(sets[np.arange(len(y)), y].mean()), "rejection": float((true_e >= 1.0 / alpha).mean()), "markov_bound": float(alpha * true_e.mean()), "avg_set_size": float(sets.sum(axis=1).mean()), "singleton_rate": float((sets.sum(axis=1) == 1).mean()), "point_accuracy": float((np.argmin(fused, axis=1) == y).mean()), } def main(): rng = np.random.default_rng(SEED) signals = np.array([1.25, 1.55, 1.85, 2.15, 2.45]) noises = np.array([1.25, 1.10, 0.95, 0.85, 0.75]) uncertainty = np.array([1.00, 0.72, 0.48, 0.28, 0.12]) ycal, ytest = rng.integers(K, size=N_CAL), rng.integers(K, size=N_TEST) cal_probs, test_probs, normalizers = [], [], [] for signal, noise in zip(signals, noises): pc = make_probs(ycal, signal, noise, rng) pt = make_probs(ytest, signal, noise, rng) raw = e_values(pc, 1.0)[np.arange(N_CAL), ycal] normalizers.append(raw.mean() * 1.01) # finite-sample safety margin cal_probs.append(pc); test_probs.append(pt) normalizers = np.asarray(normalizers) test_e = np.stack([e_values(p, z) for p, z in zip(test_probs, normalizers)], axis=2) cal_e_true = np.stack([e_values(p, z)[np.arange(N_CAL), ycal] for p, z in zip(cal_probs, normalizers)], axis=1) local_cal = cal_e_true.mean(axis=0) local_test = np.array([test_e[np.arange(N_TEST), ytest, i].mean() for i in range(M)]) out = {"seed": SEED, "n_cal": N_CAL, "n_test": N_TEST, "alpha": ALPHA, "construction": "inverse_probability_e_value", "uncertainty": uncertainty.tolist(), "normalizers": normalizers.tolist(), "local_calibration_true_e_means": local_cal.tolist(), "local_test_true_e_means": local_test.tolist(), "sweeps": {}} # Prediction 1: E[e_F] is the convex weighted mean, for all kappa. out["sweeps"]["kappa"] = [] for kappa in [0.0, 0.5, 1.0, 2.0, 5.0]: w = np.exp(-kappa * uncertainty); w /= w.sum() row = summarize(test_e, ytest, w, ALPHA) row.update(kappa=kappa, weights=w.tolist(), predicted_mean=float(w @ local_test)) row["convexity_abs_error"] = abs(row["mean_fused_true_e"] - row["predicted_mean"]) out["sweeps"]["kappa"].append(row) # Prediction 2: every active neighborhood remains valid; size can improve. out["sweeps"]["active_experts"] = [] for n in range(1, M + 1): w = np.exp(-2.0 * uncertainty[:n]); w /= w.sum() row = summarize(test_e[:, :, :n], ytest, w, ALPHA) row.update(n_active=n, weights=w.tolist(), predicted_mean=float(w @ local_test[:n])) out["sweeps"]["active_experts"].append(row) # Prediction 3: Markov rejection bound is alpha * E[e_F]. out["sweeps"]["alpha"] = [] w = np.exp(-2.0 * uncertainty); w /= w.sum() for alpha in [0.02, 0.05, 0.10, 0.20, 0.40]: row = summarize(test_e, ytest, w, alpha); row.update(alpha=alpha, weights=w.tolist()) out["sweeps"]["alpha"].append(row) best_p = test_probs[-1] best_e = e_values(best_p, normalizers[-1]) base_sets = best_e < 1.0 / ALPHA out["baseline"] = { "strongest_expert_point_accuracy": float((np.argmax(best_p, axis=1) == ytest).mean()), "strongest_expert_coverage": float(base_sets[np.arange(N_TEST), ytest].mean()), "strongest_expert_avg_set_size": float(base_sets.sum(axis=1).mean()), "unweighted": summarize(test_e, ytest, np.ones(M) / M, ALPHA), "uncertainty_weighted": summarize(test_e, ytest, w, ALPHA), } Path("results.json").write_text(json.dumps(out, indent=2)) print(json.dumps(out, indent=2)) if __name__ == "__main__": main()