Spatial-Quantile Conformal Bands for Neural Operators / run_bench.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
 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))