Distributed E-Value Prediction Sets / bench_experiment.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
 1import json, sys, random
 2from pathlib import Path
 3import numpy as np
 4import torch
 5sys.path.insert(0, "/home/maxwelhelp/all/math2nn")
 6from bench import get_dataset, make_model, train_model, evaluate, sweep_baseline, make_report
 7
 8N_TRAIN, N_TEST, EPOCHS, M = 400, 400, 3, 2
 9LR_GRID = [1e-3, 3e-3, 1e-2]
10SEEDS = tuple(range(8))
11ALPHA = 0.10
12
13def seed_all(s):
14    random.seed(s); np.random.seed(s); torch.manual_seed(s)
15    if torch.cuda.is_available(): torch.cuda.manual_seed_all(s)
16
17def train_system(seed, lr, mode, kappa=1.0):
18    seed_all(seed)
19    d = get_dataset("vision", seed, N_TRAIN, N_TEST)
20    n = len(d["xtr"]); cut = int(.8*n)
21    tr = dict(d); tr["xtr"] = d["xtr"][:cut]; tr["ytr"] = d["ytr"][:cut]
22    calx, caly = d["xtr"][cut:], d["ytr"][cut:]
23    testx, testy = d["xte"], d["yte"]
24    probs_cal, probs_test, unc = [], [], []
25    for j in range(M):
26        seed_all(seed * 100 + j)
27        net = make_model("cnn_small", d["input_shape"], d["out_dim"])
28        net, _, _ = train_model(net, tr, epochs=EPOCHS, lr=lr, batch=128, log=lambda *_: None)
29        if net is None: return float("nan"), {}
30        dev = next(net.parameters()).device
31        with torch.no_grad():
32            pc = torch.softmax(net(calx.to(dev)), 1).cpu().numpy()
33            pt = torch.softmax(net(testx.to(dev)), 1).cpu().numpy()
34        probs_cal.append(pc); probs_test.append(pt)
35        unc.append(float(-(pc * np.log(np.maximum(pc, 1e-8))).sum(1).mean()))
36    pc = np.stack(probs_cal); pt = np.stack(probs_test)
37    if mode == "baseline":
38        fused = pt.mean(0)
39        pred = fused.argmax(1)
40        return float((pred != testy.numpy()).mean()), {"accuracy": float((pred == testy.numpy()).mean())}
41    idx = np.arange(len(caly.numpy()))
42    raw = np.stack([1.0 / (10.0 * np.maximum(pc[j], 1e-8)[idx, caly.numpy()]) for j in range(M)])
43    normalizers = raw.mean(1) * 1.01
44    ev = 1.0 / (10.0 * np.maximum(pt, 1e-8)) / normalizers[:, None, None]
45    weights = np.exp(-kappa * np.asarray(unc)); weights /= weights.sum()
46    fused_e = np.sum(ev * weights[:, None, None], axis=0)
47    pred = np.argmin(fused_e, 1)
48    sets = fused_e < 1.0 / ALPHA
49    true_e = fused_e[np.arange(len(testy)), testy.numpy()]
50    return float((pred != testy.numpy()).mean()), {"accuracy": float((pred == testy.numpy()).mean()), "coverage": float(sets[np.arange(len(testy)), testy.numpy()].mean()), "avg_set_size": float(sets.sum(1).mean()), "mean_true_e": float(true_e.mean()), "predicted_mean_true_e": float(weights @ (ev[:, np.arange(len(testy)), testy.numpy()].mean(1))), "convexity_abs_error": float(abs(true_e.mean() - weights @ (ev[:, np.arange(len(testy)), testy.numpy()].mean(1))))}
51
52def main():
53    # Equal-budget sweep: every idea learning rate is also evaluated for baseline.
54    grid = [{"lr": x} for x in LR_GRID]
55    base = sweep_baseline(lambda cfg: lambda s: train_system(s, cfg["lr"], "baseline")[0], grid, seeds=(0,1,2,3))
56    best_lr = base["best_cfg"]["lr"]
57    idea_grid = [best_lr] + [x for x in LR_GRID if x != best_lr]
58    # Record the same three learning rates on the idea side before choosing one.
59    idea_sweep = [{"lr": lr, "result": evaluate(lambda s, lr=lr: train_system(s, lr, "idea", 1.0)[0], seeds=(0,1,2,3))} for lr in idea_grid]
60    idea_best_lr = min(idea_sweep, key=lambda z: z["result"]["mean"])["lr"]
61    idea = evaluate(lambda s: train_system(s, idea_best_lr, "idea", 1.0)[0], seeds=SEEDS)
62    extra = {"prediction": "fused mean e-value equals uncertainty-weighted mean local e-values", "observed": [], "confirmed": False}
63    for s in SEEDS:
64        _, z = train_system(s, best_lr, "idea", 1.0); extra["observed"].append(z)
65    errs = [z["convexity_abs_error"] for z in extra["observed"]]
66    extra["max_abs_error"] = float(max(errs)); extra["confirmed"] = bool(extra["max_abs_error"] < 1e-6)
67    base["idea_grid"] = idea_sweep
68    report = make_report("vision", "cnn_small", base, idea, extra)
69    report["idea_sweep"] = idea_sweep
70    report["notes"] = "Vision selected because labels/candidate-label e-values are classification-specific; N_TRAIN/N_TEST=400 and 3 epochs were used for budget."
71    Path("bench_report.json").write_text(json.dumps(report, indent=2))
72    print(json.dumps(report, indent=2))
73
74if __name__ == "__main__": main()