Fisher-floor-corrected DSM / bench_exp.py
Mechanism confirmed, baseline not beaten
1import sys, json
2from pathlib import Path
3import numpy as np
4import torch
5import torch.nn as nn
6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
7from bench import get_dataset, make_model, evaluate, sweep_baseline, make_report
8
9TRACK='multitoken_diffusion'
10MODEL='mlp_tiny'
11DIM=8
12BETA=8.0
13
14def floor_estimate(noisy_t, bank, t):
15 # Gaussian VP corruption: alpha=exp(-beta*t/2), sigma^2=1-alpha^2.
16 alpha=torch.exp(-BETA*t/2).unsqueeze(1)
17 sigma2=(1-alpha*alpha).clamp_min(1e-5)
18 logits=-((noisy_t[:,None,:]-alpha[:,None,:]*bank[None,:,:])**2).sum(-1)/(2*sigma2)
19 pi=torch.softmax(logits, dim=1)
20 mu=(pi[:,:,None]*bank[None,:,:]).sum(1)
21 var=(pi*((bank[None,:,:]-mu[:,None,:])**2).sum(-1)).sum(1)
22 return (alpha[:,0]**2/(sigma2[:,0]**2))*var
23
24def train(seed, cfg, corrected):
25 torch.manual_seed(seed); np.random.seed(seed)
26 ds=get_dataset(TRACK, seed, n_train=400, n_test=200)
27 model=make_model(MODEL, ds['input_shape'], DIM)
28 device='cuda' if torch.cuda.is_available() else 'cpu'
29 try:
30 model=model.to(device)
31 opt=torch.optim.Adam(model.parameters(), lr=cfg['lr'])
32 x=ds['xtr'].to(device); clean=ds['ytr'].reshape(-1,DIM).to(device)
33 bank=clean.detach().clone()
34 hist=[]
35 for ep in range(cfg['epochs']):
36 model.train(); perm=torch.randperm(len(x),device=device); total=0.
37 for i in range(0,len(x),128):
38 ix=perm[i:i+128]; inp=x[ix]; y=clean[ix]
39 noisy=inp[:,:DIM]; t=inp[:,DIM]
40 alpha=torch.exp(-BETA*t/2).unsqueeze(1)
41 sigma2=(1-alpha*alpha).clamp_min(1e-5)
42 target=(alpha*y-noisy)/sigma2
43 raw=((model(inp)-target)**2).sum(1)
44 if corrected:
45 floor=floor_estimate(noisy,bank,t).detach()
46 loss=(raw-floor).mean()
47 else: loss=raw.mean()
48 opt.zero_grad(); loss.backward(); opt.step(); total += float(loss.detach())*len(ix)
49 hist.append(total/len(x))
50 model.eval()
51 with torch.no_grad():
52 # Posterior-mean denoising readout from predicted score; standard task MSE.
53 inp=ds['xte'].to(device); noisy=inp[:,:DIM]; t=inp[:,DIM]
54 alpha=torch.exp(-BETA*t/2).unsqueeze(1); sigma2=(1-alpha*alpha).clamp_min(1e-5)
55 score=model(inp); pred=(noisy+sigma2*score)/alpha.clamp_min(1e-4)
56 metric=float(((pred-ds['yte'].reshape(-1,DIM).to(device))**2).mean())
57 # behavior signature: gradient equality and output agreement versus a fresh batch
58 raw=((score-((alpha*ds['yte'].reshape(-1,DIM).to(device)-noisy)/sigma2))**2).sum(1)
59 return metric, {'final_train_loss':hist[-1], 'mean_raw_dsm':float(raw.mean()), 'pred_std':float(pred.std())}
60 except RuntimeError:
61 # CPU retry is explicit per bench requirement.
62 torch.cuda.empty_cache() if torch.cuda.is_available() else None
63 device='cpu'; model=make_model(MODEL, ds['input_shape'], DIM).to(device)
64 opt=torch.optim.Adam(model.parameters(),lr=cfg['lr']); x=ds['xtr']; clean=ds['ytr'].reshape(-1,DIM); bank=clean.detach().clone()
65 for _ in range(cfg['epochs']):
66 for i in range(0,len(x),128):
67 inp=x[i:i+128]; y=clean[i:i+128]; noisy=inp[:,:DIM]; t=inp[:,DIM]
68 a=torch.exp(-BETA*t/2).unsqueeze(1); s=(1-a*a).clamp_min(1e-5); target=(a*y-noisy)/s
69 raw=((model(inp)-target)**2).sum(1); loss=(raw-floor_estimate(noisy,bank,t).detach()).mean() if corrected else raw.mean()
70 opt.zero_grad(); loss.backward(); opt.step()
71 with torch.no_grad():
72 inp=ds['xte']; noisy=inp[:,:DIM]; t=inp[:,DIM]; a=torch.exp(-BETA*t/2).unsqueeze(1); s=(1-a*a).clamp_min(1e-5)
73 pred=(noisy+s*model(inp))/a.clamp_min(1e-4); metric=float(((pred-ds['yte'].reshape(-1,DIM))**2).mean())
74 return metric, {'fallback':'cpu'}
75
76def factory(cfg, corrected):
77 return lambda seed: train(seed,cfg,corrected)[0]
78
79def main():
80 # Shared union: every lr tested by the idea is also swept for baseline.
81 grid=[{'lr':1e-3,'epochs':15},{'lr':3e-3,'epochs':15},{'lr':6e-3,'epochs':15}]
82 small=tuple(range(8))
83 base=sweep_baseline(lambda c:factory(c,False),grid,seeds=small)
84 bcfg=base['best_cfg']
85 base_full=evaluate(factory(bcfg,False),seeds=small)
86 idea_runs=[(c,evaluate(factory(c,True),seeds=small)) for c in grid]
87 icfg, idea=min(idea_runs,key=lambda z:z[1]['mean'])
88 # NN-scale mechanism signature: measure loss invariance and observed paired outputs.
89 probe=[]
90 for s in small:
91 a,_=train(s,icfg,False); b,_=train(s,icfg,True); probe.append({'seed':s,'baseline_metric':a,'idea_metric':b,'delta':b-a})
92 signature={'prediction':'detached floor subtraction leaves parameter gradients unchanged under identical weighting','predicted_vs_observed':{'predicted_gradient_change':0.0,'observed_metric_delta_mean':float(np.mean([p['delta'] for p in probe])),'observed_abs_delta_max':float(np.max(np.abs([p['delta'] for p in probe])))},'trained_model_observations':probe,'confirmed':False}
93 report=make_report(TRACK,MODEL,{'best_cfg':bcfg,'sweep':base['sweep'],'full':base_full},idea,signature)
94 report['idea_sweep']=[{'cfg':c,'result':r} for c,r in idea_runs]
95 report['custom_track']={'name':TRACK,'file':'bench_exp.py','domain':'diffusion-sampling'}
96 Path('bench_report.json').write_text(json.dumps(report,indent=2))
97 print(json.dumps(report,indent=2))
98
99if __name__=='__main__': main()