Square-Root Error-Density Timestep Grid / run_bench.py

Mechanism confirmed, baseline not beaten

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