Integral Sparse Dynamics Training / bench_stage2.py

Failed on benchmark

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