Integral Sparse Dynamics Training / bench_stage2.py
Failed on benchmark
1import sys, json, random
2import numpy as np
3import torch
4import torch.nn as nn
5sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
6from bench import get_dataset, make_model, train_model, evaluate, sweep_baseline, make_report
7
8EPOCHS, BATCH, NTR, NTE = 6, 128, 400, 150
9LRS = [1e-3, 3e-3, 1e-2]
10LAM = 1e-4
11
12def seed_all(s):
13 random.seed(s); np.random.seed(s); torch.manual_seed(s)
14 if torch.cuda.is_available(): torch.cuda.manual_seed_all(s)
15
16def base(cfg, seed):
17 seed_all(seed); d=get_dataset('dynamics', seed, NTR, NTE)
18 net=make_model('rnn_small', d['input_shape'], d['out_dim'])
19 _, m, _=train_model(net,d,epochs=EPOCHS,lr=cfg['lr'],batch=BATCH,
20 weight_decay=cfg['weight_decay'],log=lambda *_:None)
21 return float(m)
22
23class Net(nn.Module):
24 def __init__(self):
25 super().__init__(); self.rnn=nn.GRU(3,64,batch_first=True); self.head=nn.Linear(64,1)
26 def forward(self,x):
27 _,h=self.rnn(x.view(x.shape[0],-1,3)); return self.head(h[-1])
28 def prefixes(self,x):
29 hs,_=self.rnn(x.view(x.shape[0],-1,3)); return self.head(hs).squeeze(-1)
30
31def prox(net, amount):
32 with torch.no_grad():
33 w=net.rnn.weight_ih_l0
34 for j in range(3):
35 b=w[:,j:j+1]; n=torch.linalg.vector_norm(b)
36 b.mul_(torch.clamp(1-amount/(n+1e-12),min=0))
37
38def idea(cfg, seed, sig=False):
39 seed_all(seed); d=get_dataset('dynamics',seed,NTR,NTE)
40 dev='cpu'
41 try:
42 net=Net().to(dev); x,y=d['xtr'].to(dev),d['ytr'].to(dev); xt,yt=d['xte'].to(dev),d['yte'].to(dev)
43 opt=torch.optim.Adam(net.parameters(),lr=cfg['lr'])
44 for _ in range(EPOCHS):
45 net.train(); p=torch.randperm(len(x),device=dev)
46 for a in range(0,len(x),BATCH):
47 ix=p[a:a+BATCH]; xb,yb=x[ix],y[ix]; pp=net.prefixes(xb); th=xb.view(len(ix),8,3)[:,:,0]
48 # Integral residual over observed dt=.05 transitions, plus endpoint task loss.
49 loss=((th[:,1:]-th[:,:-1])-(pp[:,1:]-pp[:,:-1])).square().mean() + (pp[:,-1:]-yb).square().mean()
50 opt.zero_grad(); loss.backward(); opt.step(); prox(net,LAM*cfg['lr'])
51 net.eval()
52 with torch.no_grad():
53 pred=net(xt); m=float((pred-yt).square().mean().cpu())
54 pp=net.prefixes(xt); ob=xt.view(len(xt),8,3)[:,:,0]
55 z={'observed_increment_mse':float((ob[:,1:]-ob[:,:-1]).square().mean().cpu()),
56 'predicted_increment_mse':float((pp[:,1:]-pp[:,:-1]).square().mean().cpu()),
57 'endpoint_test_mse':m}
58 return (m,z) if sig else m
59 except RuntimeError:
60 raise
61
62def main():
63 # Union parity: every idea lr is in the baseline sweep; baseline's central
64 # Adam weight-decay knob is swept as well.
65 grid=[{'lr':lr,'weight_decay':wd} for lr in LRS for wd in [0.0,1e-4]]
66 base_block=sweep_baseline(lambda c: lambda s: base(c,s),grid)
67 # Same-sized three-config idea sweep, using the four sweep seeds.
68 idea_sweep=[]
69 for c in [{'lr':lr,'lam':LAM} for lr in LRS]:
70 r=evaluate(lambda s,c=c: idea(c,s),seeds=(0,1,2,3)); idea_sweep.append({'cfg':c,'mean':r['mean']})
71 best=min(idea_sweep,key=lambda z:z['mean'])['cfg']
72 idea_res=evaluate(lambda s: idea(best,s))
73 zs=[idea(best,s,True)[1] for s in range(8)]
74 sig={k:float(np.mean([z[k] for z in zs])) for k in zs[0]}
75 sig.update({'prediction':'integral residual uses observed increments instead of noisy finite-difference derivatives',
76 'confirmed':bool(sig['predicted_increment_mse'] < 2.0*sig['observed_increment_mse'])})
77 rep=make_report('dynamics','rnn_small',base_block,idea_res,sig)
78 rep['idea']['sweep']=idea_sweep; rep['idea']['best_cfg']=best
79 with open('bench_report.json','w') as f: json.dump(rep,f,indent=2)
80 print(json.dumps(rep,indent=2))
81if __name__=='__main__': main()