Second-order fusion prior for point-set diffusion / run_bench.py
Failed on benchmark
1import os, sys, json, math, importlib.util
2import numpy as np
3import torch
4import torch.nn as nn
5
6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
7from bench import make_model, make_report
8from bench.protocol import permutation_pvalue
9
10HERE = os.path.dirname(os.path.abspath(__file__))
11seeds = list(range(8))
12
13spec = importlib.util.spec_from_file_location('pointset_fusion_track', os.path.join(HERE, 'pointset_fusion_track.py'))
14track = importlib.util.module_from_spec(spec); spec.loader.exec_module(track)
15
16
17def kappa(m, beta):
18 if m * beta <= 1: raise ValueError('m beta must exceed 1')
19 return beta * beta / (8.0 * (m * beta - 1.0) * (2 * m + 1.0))
20
21
22def prior_score(x, beta=2.0, delta=0.08, correction=True, clip=20.0):
23 # x: [batch,m], centered and locally scaled before applying theorem prior.
24 m = x.shape[1]
25 scale = x.std(dim=1, keepdim=True).clamp_min(0.25)
26 a = (x - x.mean(dim=1, keepdim=True)) / scale
27 d = a[:, :, None] - a[:, None, :]
28 mask = 1.0 - torch.eye(m, device=x.device)[None]
29 rep = beta * (d / (d*d + delta*delta) * mask).sum(dim=2)
30 corr = -2.0 * kappa(m, beta) * (d * mask).sum(dim=2) if correction else 0.0
31 q = rep + corr
32 return q.clamp(-clip, clip)
33
34
35def train_one(seed, cfg, idea):
36 torch.manual_seed(seed); np.random.seed(seed)
37 d0 = track.get_dataset(seed, n_train=400, n_test=160)
38 xtr = torch.tensor(d0['xtr']); ytr = torch.tensor(d0['ytr'])
39 xte = torch.tensor(d0['xte']); yte = torch.tensor(d0['yte'])
40 net = make_model('mlp_tiny', (6,), 6)
41 dev = 'cuda' if torch.cuda.is_available() else 'cpu'
42 try:
43 net = net.to(dev); xtr=xtr.to(dev); ytr=ytr.to(dev); xte=xte.to(dev); yte=yte.to(dev)
44 opt = torch.optim.Adam(net.parameters(), lr=cfg['lr'], weight_decay=cfg['weight_decay'])
45 gen = torch.Generator(device=dev).manual_seed(seed + 991)
46 bs=128
47 for _ in range(18):
48 net.train(); perm=torch.randperm(len(xtr), generator=gen, device=dev)
49 for start in range(0, len(xtr), bs):
50 ix=perm[start:start+bs]; clean=ytr[ix]; noisy=xtr[ix]
51 t=torch.rand((len(ix),1), generator=gen, device=dev)
52 sigma=0.08 + 0.34*t; alpha=torch.sqrt(1.0-sigma*sigma)
53 xt=alpha*noisy + sigma*torch.randn(xt_shape := noisy.shape, generator=gen, device=dev)
54 # Network predicts the diffusion score; denoising estimate is the
55 # standard Tweedie reconstruction used by both systems.
56 pred=net(xt)
57 target=(clean-xt)/(sigma*sigma)
58 loss=((pred-target)**2).mean()
59 if idea:
60 # Local score prior evaluated on the model's own denoised
61 # point-set estimate, not on an oracle target.
62 xhat=xt + sigma*sigma*pred
63 q=prior_score(xhat, beta=cfg['beta'], delta=cfg['delta'], correction=True)
64 loss=loss + cfg['lam'] * ((pred-q)**2).mean() * (sigma < 0.55).float().mean()
65 opt.zero_grad(); loss.backward(); torch.nn.utils.clip_grad_norm_(net.parameters(), 10.0); opt.step()
66 net.eval()
67 with torch.no_grad():
68 # Fixed moderate-noise test score, standard denoising MSE.
69 sigma=torch.full((len(xte),1), 0.25, device=dev)
70 xt=xte + sigma*torch.randn(xte.shape, generator=gen, device=dev)
71 pred=net(xt); xhat=xt + sigma*sigma*pred
72 metric=((xhat-yte)**2).mean().item()
73 # Behavioral signature: normalized short-gap collision rate of the
74 # trained system's predictions, measured identically for both sides.
75 ss=torch.sort(xhat, dim=1).values
76 gaps=ss[:,1:]-ss[:,:-1]; gaps=gaps/(gaps.mean(dim=1,keepdim=True)+1e-6)
77 collision=(gaps < 0.10).float().mean().item()
78 mean_gap=gaps.mean().item()
79 return {'metric': float(metric), 'collision': float(collision), 'mean_gap': float(mean_gap)}
80 except RuntimeError:
81 # Robust CPU fallback for shared/unsupported CUDA environments.
82 if dev == 'cuda':
83 torch.cuda.empty_cache()
84 old=torch.cuda.is_available
85 torch.cuda.is_available=lambda: False
86 try: return train_one(seed, cfg, idea)
87 finally: torch.cuda.is_available=old
88 raise
89
90
91def eval_cfg(cfg, idea):
92 vals=[]; cols=[]; gaps=[]
93 for s in seeds:
94 r=train_one(s,cfg,idea); vals.append(r['metric']); cols.append(r['collision']); gaps.append(r['mean_gap'])
95 return {'mean':float(np.mean(vals)), 'std':float(np.std(vals)), 'per_seed':vals, 'n':len(vals),
96 'collision_per_seed':cols, 'collision_mean':float(np.mean(cols)), 'mean_gap':float(np.mean(gaps)), 'cfg':cfg}
97
98
99def main():
100 # Union of all learning rates appears on both sides; weight decay is the
101 # baseline's central optimizer knob and is swept equally for every lr.
102 lrs=[1e-3, 2e-3, 3e-3]; wds=[0.0, 1e-4]
103 base_grid=[{'lr':lr,'weight_decay':wd} for lr in lrs for wd in wds]
104 base_sweep=[]
105 for cfg in base_grid:
106 r=eval_cfg(cfg,False); base_sweep.append({'cfg':cfg,'mean':r['mean']})
107 best=min(base_sweep,key=lambda z:z['mean'])['cfg']
108 base_full=eval_cfg(best,False)
109 # Three idea settings: best baseline lr plus two nearby parity lrs.
110 idea_grid=[{'lr':lr,'weight_decay':best['weight_decay'],'lam':lam,'beta':2.0,'delta':0.08}
111 for lr,lam in [(best['lr'],0.10),(2e-3,0.10),(1e-3,0.10)]]
112 idea_results=[eval_cfg(c,True) for c in idea_grid]
113 idea=min(idea_results,key=lambda z:z['mean'])
114 diffs=[a-b for a,b in zip(idea['per_seed'],base_full['per_seed'])]
115 cmp={'delta_mean':float(np.mean(diffs)), 'idea_wins':sum(x<0 for x in diffs), 'n_pairs':8,
116 'per_seed_diffs':[float(x) for x in diffs], 'p_value':float(permutation_pvalue(diffs))}
117 if cmp['delta_mean']<0 and cmp['p_value']<0.05: cmp['verdict']='idea better (significant)'; cmp['system_worked']=True
118 elif cmp['delta_mean']>0 and cmp['p_value']<0.05: cmp['verdict']='idea worse (significant)'; cmp['system_worked']=False
119 else: cmp['verdict']='no significant win'; cmp['system_worked']=False
120 # Retest stage-1 prediction at NN scale: prior should reduce predicted
121 # short-gap collisions relative to baseline. This is not the primary metric.
122 pred_delta=idea['collision_mean']-base_full['collision_mean']
123 observed_delta=float(np.mean(np.array(idea['per_seed'])-np.array(base_full['per_seed'])))
124 sig={'quantity':'normalized predicted short-gap collision rate',
125 'baseline_predicted_collision':base_full['collision_mean'],
126 'idea_predicted_collision':idea['collision_mean'],
127 'predicted_delta':float(pred_delta),
128 'observed_test_mse_delta':observed_delta,
129 'confirmed':bool(pred_delta < 0)}
130 base_block={'best_cfg':best,'sweep':base_sweep,'full':base_full}
131 report=make_report('unordered_pointset_denoising','mlp_tiny',base_block,idea,
132 {'custom_track':{'name':'unordered_pointset_denoising','file':'pointset_fusion_track.py','domain':'point-set-diffusion'},
133 'idea_settings':idea_results,'mechanism_signature':sig})
134 report['comparison']=cmp
135 with open(os.path.join(HERE,'bench_report.json'),'w') as f: json.dump(report,f,indent=2)
136 print(json.dumps(report,indent=2))
137
138if __name__=='__main__': main()