Second-order fusion prior for point-set diffusion / run_bench.py

Failed on benchmark

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