Spatial-Quantile Conformal Bands for Neural Operators / run_bench.py
Mechanism confirmed, baseline not beaten
1import sys, json, math, importlib.util
2from pathlib import Path
3import numpy as np
4import torch
5sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
6from bench import make_model, train_model, evaluate, sweep_baseline, make_report
7HERE = Path(__file__).resolve().parent
8sp = importlib.util.spec_from_file_location('poisson_field_track', HERE/'poisson_field_track.py')
9track = importlib.util.module_from_spec(sp); sp.loader.exec_module(track)
10
11def wquant(v, w, mass):
12 o = np.argsort(v, kind='mergesort'); c = np.cumsum(w[o]/w.sum())
13 return float(v[o][np.searchsorted(c, mass, side='left')])
14
15def cq(scores, alpha=.1):
16 a = np.sort(np.asarray(scores)); k = min(len(a), int(math.ceil((len(a)+1)*(1-alpha))))
17 return float(a[k-1])
18
19def run_once(seed, lr, mode, gamma=.1, collect=False):
20 raw = track.get_dataset(seed, 400, 200)
21 xfit,yfit = raw['xtr'][:280],raw['ytr'][:280]
22 xcal,ycal = raw['xtr'][280:],raw['ytr'][280:]
23 ds={'xtr':torch.tensor(xfit),'ytr':torch.tensor(yfit),'xte':torch.tensor(raw['xte']),'yte':torch.tensor(raw['yte']), 'task':'regression','metric':'mse','input_shape':(6,),'out_dim':64}
24 torch.manual_seed(10000+seed); np.random.seed(10000+seed)
25 net=make_model('mlp_tiny',ds['input_shape'],64)
26 net,_,_=train_model(net,ds,epochs=18,lr=lr,batch=128)
27 device=next(net.parameters()).device
28 with torch.no_grad():
29 pc=net(torch.tensor(xcal, device=device)).cpu().numpy(); pt=net(torch.tensor(raw['xte'], device=device)).cpu().numpy()
30 rc=np.abs(ycal-pc); rt=np.abs(raw['yte']-pt)
31 scale=np.maximum(np.quantile(np.abs(yfit-net(torch.tensor(xfit, device=device)).detach().cpu().numpy()),.75,axis=0),.01)
32 normc=rc/scale; normt=rt/scale
33 weights=np.ones(normc.shape[1])/normc.shape[1]
34 if mode=='max': scores=normc.max(axis=1); q=cq(scores)
35 else:
36 scores=np.array([wquant(row,weights,1-gamma) for row in normc]); q=cq(scores)
37 frac=(normt<=q).mean(axis=1)
38 mse=float(np.mean((pt-raw['yte'])**2))
39 out={'mse':mse,'q':q,'mean_width':float(2*q*scale.mean()),'event_rate':float(np.mean(frac>=1-gamma)),'mean_fraction':float(frac.mean())}
40 return out if collect else mse
41
42if __name__ == '__main__':
43 lrs = [0.0015, 0.003, 0.006]
44 seeds = tuple(range(8))
45 base = sweep_baseline(
46 lambda cfg: (lambda s: run_once(s, cfg['lr'], 'max')),
47 [{'lr': lr} for lr in lrs], seeds=(0,1,2,3))
48 # Idea uses the same three learning rates: search-space parity is exact.
49 idea_runs = []
50 for lr in lrs:
51 r = evaluate(lambda s: run_once(s, lr, 'quantile', gamma=.10), seeds=seeds)
52 idea_runs.append({'cfg': {'lr': lr, 'gamma': .10}, 'result': r})
53 best = min(idea_runs, key=lambda z: z['result']['mean'])
54 idea = best['result']
55 # Measure the trained-model mechanism on all paired test predictions.
56 sig = []
57 for s in seeds:
58 b = run_once(s, base['best_cfg']['lr'], 'max', collect=True)
59 q = run_once(s, best['cfg']['lr'], 'quantile', gamma=.10, collect=True)
60 sig.append({'seed': s, 'max_width': b['mean_width'], 'quantile_width': q['mean_width'],
61 'max_event_rate': b['event_rate'], 'quantile_event_rate': q['event_rate']})
62 ms = {
63 'prediction': 'spatial quantile should reduce mean band width for gamma=.10 while retaining approximately 90% domain-event rate',
64 'observed_mean_max_width': float(np.mean([x['max_width'] for x in sig])),
65 'observed_mean_quantile_width': float(np.mean([x['quantile_width'] for x in sig])),
66 'observed_mean_max_event_rate': float(np.mean([x['max_event_rate'] for x in sig])),
67 'observed_mean_quantile_event_rate': float(np.mean([x['quantile_event_rate'] for x in sig])),
68 'confirmed': bool(np.mean([x['quantile_width'] for x in sig]) < np.mean([x['max_width'] for x in sig]) and np.mean([x['quantile_event_rate'] for x in sig]) >= .85),
69 'per_seed': sig,
70 }
71 report = make_report('poisson_field_operator', 'mlp_tiny', base, idea,
72 {'mechanism_signature': ms, 'idea_sweep': idea_runs,
73 'custom_track': {'name': 'poisson_field_operator', 'file': 'poisson_field_track.py', 'domain': 'pde'}})
74 Path('bench_report.json').write_text(json.dumps(report, indent=2))
75 print(json.dumps(report, indent=2))