Square-Root Error-Density Timestep Grid / run_bench.py
Mechanism confirmed, baseline not beaten
1import os, sys, json, math
2import numpy as np
3import torch
4import torch.nn as nn
5sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
6from bench import evaluate, sweep_baseline, make_report, permutation_pvalue, get_dataset as bench_get_dataset
7from custom_diffusion_track import META
8
9SEEDS = tuple(range(8))
10DEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'
11
12class ScoreNet(nn.Module):
13 def __init__(self):
14 super().__init__()
15 self.net = nn.Sequential(nn.Linear(4+8+1,64), nn.SiLU(), nn.Linear(64,64), nn.SiLU(), nn.Linear(64,8))
16 def forward(self, c, x, t):
17 return self.net(torch.cat([c,x,t],1))
18
19def fit(seed, lr, epochs=12):
20 torch.manual_seed(seed); np.random.seed(seed)
21 d = bench_get_dataset('conditional_sequence_diffusion', seed, 400, 100)
22 c = torch.tensor(d['xtr'], dtype=torch.float32)
23 y = torch.tensor(d['ytr'], dtype=torch.float32).reshape(-1, 8)
24 net = ScoreNet()
25 try:
26 dev = torch.device(DEVICE); net.to(dev); c,y=c.to(dev),y.to(dev)
27 opt = torch.optim.Adam(net.parameters(), lr=lr)
28 gen = torch.Generator(device=dev).manual_seed(seed+991)
29 for ep in range(epochs):
30 perm = torch.randperm(len(c), generator=gen, device=dev)
31 for ii in range(0,len(c),128):
32 ix=perm[ii:ii+128]; t=torch.rand((len(ix),1),generator=gen,device=dev)
33 noise=torch.randn((len(ix),8),generator=gen,device=dev)
34 xt=torch.sqrt(1-t)*y[ix] + torch.sqrt(t)*noise
35 pred=net(c[ix],xt,t)
36 loss=((pred-y[ix])**2).mean()
37 opt.zero_grad(); loss.backward(); opt.step()
38 return net, d
39 except RuntimeError:
40 net=ScoreNet(); opt=torch.optim.Adam(net.parameters(),lr=lr)
41 for ep in range(epochs):
42 ix=torch.randperm(len(c), device='cpu')
43 for ii in range(0,len(c),128):
44 j=ix[ii:ii+128]; t=torch.rand(len(j),1); noise=torch.randn(len(j),8)
45 xt=torch.sqrt(1-t)*y[j]+torch.sqrt(t)*noise
46 loss=((net(c[j],xt,t)-y[j])**2).mean(); opt.zero_grad(); loss.backward(); opt.step()
47 return net,d
48
49def density_grid(net,d,n,seed):
50 dev=next(net.parameters()).device
51 c=torch.tensor(d['xte'][:64],dtype=torch.float32,device=dev); y=torch.tensor(d['yte'],dtype=torch.float32,device=dev).reshape(-1,8)[:64]
52 rng=torch.Generator(device=dev).manual_seed(seed+2222); noise=torch.randn(y.shape,generator=rng,device=dev)
53 ts=torch.linspace(.01,.99,33,device=dev); vals=[]
54 with torch.no_grad():
55 for t in ts:
56 x=torch.sqrt(1-t)*y+torch.sqrt(t)*noise
57 h=1e-3; t2=torch.clamp(t+h,max=.999)
58 p1=net(c,x,t.expand(len(c),1)); p2=net(c,x,t2.expand(len(c),1))
59 vals.append(float(((p2-p1)/h).pow(2).mean().cpu()))
60 a=np.maximum(np.asarray(vals),1e-8); w=np.sqrt(a); cum=np.r_[0,np.cumsum((w[1:]+w[:-1])*np.diff(ts.cpu().numpy())/2)]
61 targets=np.linspace(cum[0],cum[-1],n+1); return np.interp(targets,cum,ts.cpu().numpy()), a, ts.cpu().numpy()
62
63def sample_metric(net,d,n,adaptive,seed, collect=False):
64 dev=next(net.parameters()).device; c=torch.tensor(d['xte'],dtype=torch.float32,device=dev)
65 gen=torch.Generator(device=dev).manual_seed(seed+7788); x=torch.randn((len(c),8),generator=gen,device=dev)
66 if adaptive: grid,a,ts=density_grid(net,d,n,seed)
67 else: grid=np.linspace(.01,.99,n+1); a=None; ts=None
68 net.eval()
69 with torch.no_grad():
70 for hi,lo in zip(grid[::-1][:-1],grid[::-1][1:]):
71 t=torch.full((len(c),1),float(hi),device=dev)
72 x0=net(c,x,t); step=(hi-lo)/max(hi,1e-6); x=(1-step)*x+step*x0
73 mse=float(((x-torch.tensor(d['yte'],dtype=torch.float32,device=dev).reshape(-1, 8))**2).mean().cpu())
74 if collect: return mse,grid,a,ts
75 return mse
76
77def train_eval(seed,cfg,adaptive):
78 net,d=fit(seed,float(cfg['lr']),int(cfg['epochs']))
79 return sample_metric(net,d,int(cfg['steps']),adaptive,seed)
80
81def make_fn(cfg,adaptive):
82 return lambda seed: train_eval(seed,cfg,adaptive)
83
84def main():
85 # Same union of method hyperparameters is tested on both sides.
86 grid=[{'lr':x,'epochs':12,'steps':n} for x in (0.001,0.003,0.01) for n in (8,16)]
87 base=sweep_baseline(lambda cfg: make_fn(cfg,False),grid,seeds=SEEDS)
88 # idea is evaluated at every baseline candidate, then best selected fairly
89 idea_trials=[]
90 for cfg in grid:
91 r=evaluate(make_fn(cfg,True),SEEDS); idea_trials.append({'cfg':cfg,'mean':r['mean'],'full':r})
92 best=min(idea_trials,key=lambda z:z['mean']); idea=best['full']; idea_cfg=best['cfg']
93 rep=make_report(META['name'],'score_mlp',{'best_cfg':base['best_cfg'],'sweep':base['sweep'],'full':base['full']},idea,extra={})
94 # measured signature from one trained benchmark model: allocation vs temporal density,
95 # and predicted equalization vs observed one-step proxy on the same learned net.
96 net,d=fit(0,float(idea_cfg['lr']),int(idea_cfg['epochs']))
97 _,g,a,ts=sample_metric(net,d,int(idea_cfg['steps']),True,0,True)
98 widths=np.diff(g); amid=np.interp((g[:-1]+g[1:])/2,ts,a)
99 corr=float(np.corrcoef(np.log(widths),np.log(1/np.sqrt(amid)))[0,1])
100 uniform=np.linspace(.01,.99,int(idea_cfg['steps'])+1)
101 # local model-change proxy: squared temporal change times interval^2
102 costs=np.interp((g[:-1]+g[1:])/2,ts,a)*widths**2
103 rep['mechanism_signature']={'density_mean':float(a.mean()),'density_max':float(a.max()),'width_inverse_sqrt_corr':corr,'observed_cost_cv':float(np.std(costs)/np.mean(costs)),'predicted_cost_equalization':'low CV under sqrt allocation','confirmed':bool(corr>.8 and np.std(costs)/np.mean(costs)<.5)}
104 rep['idea']['trials']=[{'cfg': z['cfg'], 'mean': z['mean']} for z in idea_trials]
105 rep['custom_track']={'name':META['name'],'file':'custom_diffusion_track.py','domain':META['domain']}
106 rep['comparison']['permutation_pvalue']=permutation_pvalue([i-j for i,j in zip(rep['idea']['per_seed'],rep['baseline']['full']['per_seed'])])
107 with open('bench_report.json','w') as f: json.dump(rep,f,indent=2)
108 print(json.dumps(rep,indent=2))
109if __name__=='__main__': main()