import sys, json, random import numpy as np import torch import torch.nn as nn sys.path.insert(0, '/home/maxwelhelp/all/math2nn') from bench import get_dataset, make_model, train_model, evaluate, sweep_baseline, make_report EPOCHS, BATCH, NTR, NTE = 6, 128, 400, 150 LRS = [1e-3, 3e-3, 1e-2] LAM = 1e-4 def seed_all(s): random.seed(s); np.random.seed(s); torch.manual_seed(s) if torch.cuda.is_available(): torch.cuda.manual_seed_all(s) def base(cfg, seed): seed_all(seed); d=get_dataset('dynamics', seed, NTR, NTE) net=make_model('rnn_small', d['input_shape'], d['out_dim']) _, m, _=train_model(net,d,epochs=EPOCHS,lr=cfg['lr'],batch=BATCH, weight_decay=cfg['weight_decay'],log=lambda *_:None) return float(m) class Net(nn.Module): def __init__(self): super().__init__(); self.rnn=nn.GRU(3,64,batch_first=True); self.head=nn.Linear(64,1) def forward(self,x): _,h=self.rnn(x.view(x.shape[0],-1,3)); return self.head(h[-1]) def prefixes(self,x): hs,_=self.rnn(x.view(x.shape[0],-1,3)); return self.head(hs).squeeze(-1) def prox(net, amount): with torch.no_grad(): w=net.rnn.weight_ih_l0 for j in range(3): b=w[:,j:j+1]; n=torch.linalg.vector_norm(b) b.mul_(torch.clamp(1-amount/(n+1e-12),min=0)) def idea(cfg, seed, sig=False): seed_all(seed); d=get_dataset('dynamics',seed,NTR,NTE) dev='cpu' try: net=Net().to(dev); x,y=d['xtr'].to(dev),d['ytr'].to(dev); xt,yt=d['xte'].to(dev),d['yte'].to(dev) opt=torch.optim.Adam(net.parameters(),lr=cfg['lr']) for _ in range(EPOCHS): net.train(); p=torch.randperm(len(x),device=dev) for a in range(0,len(x),BATCH): ix=p[a:a+BATCH]; xb,yb=x[ix],y[ix]; pp=net.prefixes(xb); th=xb.view(len(ix),8,3)[:,:,0] # Integral residual over observed dt=.05 transitions, plus endpoint task loss. loss=((th[:,1:]-th[:,:-1])-(pp[:,1:]-pp[:,:-1])).square().mean() + (pp[:,-1:]-yb).square().mean() opt.zero_grad(); loss.backward(); opt.step(); prox(net,LAM*cfg['lr']) net.eval() with torch.no_grad(): pred=net(xt); m=float((pred-yt).square().mean().cpu()) pp=net.prefixes(xt); ob=xt.view(len(xt),8,3)[:,:,0] z={'observed_increment_mse':float((ob[:,1:]-ob[:,:-1]).square().mean().cpu()), 'predicted_increment_mse':float((pp[:,1:]-pp[:,:-1]).square().mean().cpu()), 'endpoint_test_mse':m} return (m,z) if sig else m except RuntimeError: raise def main(): # Union parity: every idea lr is in the baseline sweep; baseline's central # Adam weight-decay knob is swept as well. grid=[{'lr':lr,'weight_decay':wd} for lr in LRS for wd in [0.0,1e-4]] base_block=sweep_baseline(lambda c: lambda s: base(c,s),grid) # Same-sized three-config idea sweep, using the four sweep seeds. idea_sweep=[] for c in [{'lr':lr,'lam':LAM} for lr in LRS]: r=evaluate(lambda s,c=c: idea(c,s),seeds=(0,1,2,3)); idea_sweep.append({'cfg':c,'mean':r['mean']}) best=min(idea_sweep,key=lambda z:z['mean'])['cfg'] idea_res=evaluate(lambda s: idea(best,s)) zs=[idea(best,s,True)[1] for s in range(8)] sig={k:float(np.mean([z[k] for z in zs])) for k in zs[0]} sig.update({'prediction':'integral residual uses observed increments instead of noisy finite-difference derivatives', 'confirmed':bool(sig['predicted_increment_mse'] < 2.0*sig['observed_increment_mse'])}) rep=make_report('dynamics','rnn_small',base_block,idea_res,sig) rep['idea']['sweep']=idea_sweep; rep['idea']['best_cfg']=best with open('bench_report.json','w') as f: json.dump(rep,f,indent=2) print(json.dumps(rep,indent=2)) if __name__=='__main__': main()