Finite-Width NNGP Covariance Stabilizer / stage2_bench.py
Failed on benchmark
1import json, math, random, sys
2import numpy as np
3import torch
4from torch import nn
5sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
6import bench
7
8SEEDS = tuple(range(8))
9TRACK='sequence'; MODEL='transformer_tiny'; NTR=400; NTE=200
10# Shared search space: every idea learning rate is also a baseline configuration.
11LR_GRID=[0.001,0.003,0.009]
12EPOCHS=12; BATCH=64
13
14def seed_all(s):
15 random.seed(s); np.random.seed(s); torch.manual_seed(s)
16 if torch.cuda.is_available(): torch.cuda.manual_seed_all(s)
17
18def make(seed):
19 seed_all(seed); ds=bench.get_dataset(TRACK, seed, NTR, NTE)
20 return ds, bench.make_model(MODEL, ds['input_shape'], ds['out_dim'])
21
22def target_cov(x, model):
23 # NNGP for the first shared linear projection, pooled over sequence positions.
24 # PyTorch Linear(1,d) has init variance 1/3; positional vectors are shared.
25 c=1.0/3.0
26 z=x.mean(1)
27 K=c*(z[:,None]*z[None,:])
28 p=model.pos[:, :x.shape[1], :].mean(1)
29 K=K + (p @ p.T)/p.shape[-1]
30 return K.detach()
31
32def pooled_cov(h):
33 q=h.mean(1)
34 return q @ q.T / q.shape[-1]
35
36def train_idea(seed, lr, lam, return_model=False):
37 ds, net=make(seed)
38 x,y=ds['xtr'],ds['ytr']; lossf=nn.MSELoss()
39 # Robust CUDA fallback, matching the benchmark's allowed device policy.
40 devices=['cuda','cpu'] if torch.cuda.is_available() else ['cpu']
41 last=None
42 for dev in devices:
43 try:
44 net=net.to(dev); xx,yy=x.to(dev),y.to(dev)
45 opt=torch.optim.Adam(net.parameters(),lr=lr)
46 for ep in range(EPOCHS):
47 net.train(); perm=torch.randperm(len(xx),device=dev)
48 for i in range(0,len(xx),BATCH):
49 ix=perm[i:i+BATCH]; xb,yb=xx[ix],yy[ix]
50 h=net.inp(xb.unsqueeze(-1))+net.pos[:,:xb.shape[1]]
51 pred=net.head(net.enc(h).reshape(xb.shape[0],-1))
52 task=lossf(pred,yb)
53 cov=((pooled_cov(h)-target_cov(xb,net))**2).mean()
54 loss=task+lam*cov
55 opt.zero_grad(); loss.backward(); opt.step()
56 net.eval()
57 with torch.no_grad():
58 pred=net(ds['xte'].to(dev)); metric=float(((pred-ds['yte'].to(dev))**2).mean())
59 if return_model: return net, metric, ds, dev
60 return metric
61 except RuntimeError as e:
62 last=e
63 if dev=='cuda':
64 net=net.to('cpu'); continue
65 raise
66 raise last
67
68def train_base(seed, lr, return_model=False):
69 ds,net=make(seed)
70 out=bench.train_model(net,ds,epochs=EPOCHS,lr=lr,batch=BATCH,weight_decay=0.0,log=lambda _:None)
71 if return_model: return out[0],float(out[1]),ds
72 return float(out[1])
73
74def signature(base_cfg, idea_cfg):
75 bd=[]; idd=[]
76 for s in SEEDS:
77 bm,bmse,ds=train_base(s,base_cfg['lr'],True)
78 im,imse,ids,dev=train_idea(s,idea_cfg['lr'],idea_cfg['lambda_cov'],True)
79 with torch.no_grad():
80 xb=ds['xte'][:64].to(next(bm.parameters()).device)
81 hb=bm.inp(xb.unsqueeze(-1))+bm.pos[:,:xb.shape[1]]
82 kb=target_cov(xb,bm).to(hb.device)
83 xi=ids['xte'][:64].to(dev); hi=im.inp(xi.unsqueeze(-1))+im.pos[:,:xi.shape[1]]
84 ki=target_cov(xi,im).to(hi.device)
85 bd.append(float(((pooled_cov(hb)-kb)**2).mean()))
86 idd.append(float(((pooled_cov(hi)-ki)**2).mean()))
87 predicted=1/math.sqrt(64)
88 observed_ratio=float(np.mean(idd)/np.mean(bd))
89 return {'quantity':'pooled input-projection covariance deviation on held-out benchmark windows',
90 'predicted_finite_width_scale_at_width_64':predicted,
91 'observed_baseline_mean':float(np.mean(bd)),
92 'observed_idea_mean':float(np.mean(idd)),
93 'observed_idea_over_baseline':observed_ratio,
94 'confirmed': bool(observed_ratio < 0.9),
95 'note':'The O(n^-1/2) prediction is a scale law, not an absolute covariance value; this signature tests reduction on trained systems.'}
96
97def main():
98 base_grid=[{'lr':v} for v in LR_GRID]
99 base=bench.sweep_baseline(lambda c: lambda s: train_base(s,c['lr']),base_grid,seeds=SEEDS)
100 idea_trials=[]
101 for cfg in [{'lr':base['best_cfg']['lr'],'lambda_cov':1e-4}, {'lr':0.001,'lambda_cov':1e-3}, {'lr':0.009,'lambda_cov':1e-2}]:
102 r=bench.evaluate(lambda s,c=cfg: train_idea(s,c['lr'],c['lambda_cov']),SEEDS)
103 idea_trials.append({'cfg':cfg,'result':r})
104 best=min(idea_trials,key=lambda z:z['result']['mean'])
105 rep=bench.make_report(TRACK,MODEL,base,best['result'],{'mechanism_signature':signature(base['best_cfg'],best['cfg']), 'idea_sweep':idea_trials, 'epochs':EPOCHS,'batch':BATCH})
106 with open('bench_report.json','w') as f: json.dump(rep,f,indent=2)
107 print(json.dumps(rep,indent=2))
108if __name__=='__main__': main()