import sys,json,random from pathlib import Path import numpy as np, torch import torch.nn as nn sys.path.insert(0,'/home/maxwelhelp/all/math2nn') from bench import get_dataset,train_model,sweep_baseline,make_report,count_params SEEDS=tuple(range(8)); SWEEP_SEEDS=tuple(range(4)); EPOCHS=8; BATCH=128; HIDDEN=60 class VanillaRNN(nn.Module): def __init__(self,out_dim=1): super().__init__(); self.inp=nn.Linear(3,HIDDEN); self.cell=nn.GRUCell(HIDDEN,HIDDEN); self.head=nn.Linear(HIDDEN,out_dim) def forward(self,x): s=x.view(x.shape[0],-1,3); h=torch.tanh(self.inp(s[:,0])) for k in range(s.shape[1]): h=self.cell(self.inp(s[:,k]),h) return self.head(h) class LiePoissonRNN(nn.Module): def __init__(self,out_dim=1,use_coadjoint=True): super().__init__(); self.use_coadjoint=use_coadjoint; self.inp=nn.Linear(3,HIDDEN) self.hamiltonian=nn.Sequential(nn.Linear(HIDDEN,96),nn.Tanh(),nn.Linear(96,1)); self.head=nn.Linear(HIDDEN,out_dim) def vector_field(self,q,p,u): with torch.enable_grad(): qq=q.requires_grad_(True); pp=p.requires_grad_(True) H=self.hamiltonian(torch.cat([qq,pp],1)).sum() gq,gp=torch.autograd.grad(H,(qq,pp),create_graph=self.training,retain_graph=self.training) qdot=gp; pdot=-gq if self.use_coadjoint: q3=q.reshape(-1,10,3); p3=p.reshape(-1,10,3); v3=qdot.reshape(-1,10,3); pd=pdot.reshape(-1,10,3).clone() pd[:,:,0]+=p3[:,:,1]*v3[:,:,2]-p3[:,:,2]*v3[:,:,1] pd[:,:,1]+=p3[:,:,2]*v3[:,:,0]-p3[:,:,0]*v3[:,:,2] pd[:,:,2]+=p3[:,:,0]*v3[:,:,1]-p3[:,:,1]*v3[:,:,0]; pdot=pd.reshape(-1,HIDDEN//2) return qdot,pdot def transition(self,q,p,u): qd,pd=self.vector_field(q,p,u); return q+qd,p+pd def forward(self,x): s=x.view(x.shape[0],-1,3); q,p=torch.chunk(torch.tanh(self.inp(s[:,0])),2,1) for k in range(s.shape[1]): q,p=self.transition(q,p,self.inp(s[:,k])); q,p=torch.tanh(q),torch.tanh(p) return self.head(torch.cat([q,p],1)) 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 train(kind,cfg,seed,collect=False): seed_all(seed); d=get_dataset('dynamics',seed,n_train=400,n_test=200); m=VanillaRNN() if kind=='base' else LiePoissonRNN() net,metric,_=train_model(m,d,epochs=EPOCHS,lr=cfg['lr'],batch=BATCH,log=lambda *a:None) if net is None:return (float('nan'),[],0) if collect else float('nan') if not collect:return float(metric) net.eval(); dev=next(net.parameters()).device; x=d['xte'][:64].to(dev); s=x.view(x.shape[0],-1,3); q,p=torch.chunk(torch.tanh(net.inp(s[:,0])),2,1); vals=[] with torch.enable_grad(): for a in [.25,.5,1.]: net.use_coadjoint=True; z=net.transition(q*a,p*a,net.inp(s[:,0])); net.use_coadjoint=False; w=net.transition(q*a,p*a,net.inp(s[:,0])); vals.append(float(torch.sqrt(torch.mean((z[1]-w[1])**2)).detach())) return float(metric),vals,count_params(net) def ev(v): a=np.asarray(v,float);return {'mean':float(np.nanmean(a)),'std':float(np.nanstd(a)),'per_seed':a.tolist(),'n':int(np.isfinite(a).sum())} def main(): grid=[{'lr':1e-3},{'lr':3e-3},{'lr':1e-2}]; base=sweep_baseline(lambda c:lambda s:train('base',c,s),grid,seeds=SWEEP_SEEDS); runs=[] for c in grid:runs.append({'cfg':c,'eval':ev([train('idea',c,s) for s in SEEDS])}) best=min(runs,key=lambda r:r['eval']['mean']); idea=best['eval']; sig=[train('idea',best['cfg'],s,True) for s in SEEDS[:4]]; obs=np.mean([x[1] for x in sig],0); slope=float(np.polyfit(np.log([.25,.5,1.]),np.log(np.maximum(obs,1e-12)),1)[0]) extra={'mechanism_signature':{'trained_model':'LiePoissonRNN','scales':[.25,.5,1.],'observed_rms_effect':obs.tolist(),'observed_loglog_exponent':slope,'confirmed':bool(1.5<=slope<=2.5)},'idea_sweep':runs,'parameter_count_baseline':count_params(VanillaRNN()),'parameter_count_idea':count_params(LiePoissonRNN())} rep=make_report('dynamics','rnn_small',{'best_cfg':base['best_cfg'],'sweep':base['sweep'],'full':base['full']},idea,extra);Path('bench_report.json').write_text(json.dumps(rep,indent=2));print(json.dumps(rep,indent=2)) if __name__=='__main__':main()