Lie-Poisson Hamiltonian latent block / bench_run.py

Failed on benchmark

Raw ⬇ ZIP
 1import sys,json,random
 2from pathlib import Path
 3import numpy as np, torch
 4import torch.nn as nn
 5sys.path.insert(0,'/home/maxwelhelp/all/math2nn')
 6from bench import get_dataset,train_model,sweep_baseline,make_report,count_params
 7SEEDS=tuple(range(8)); SWEEP_SEEDS=tuple(range(4)); EPOCHS=8; BATCH=128; HIDDEN=60
 8class VanillaRNN(nn.Module):
 9 def __init__(self,out_dim=1):
10  super().__init__(); self.inp=nn.Linear(3,HIDDEN); self.cell=nn.GRUCell(HIDDEN,HIDDEN); self.head=nn.Linear(HIDDEN,out_dim)
11 def forward(self,x):
12  s=x.view(x.shape[0],-1,3); h=torch.tanh(self.inp(s[:,0]))
13  for k in range(s.shape[1]): h=self.cell(self.inp(s[:,k]),h)
14  return self.head(h)
15class LiePoissonRNN(nn.Module):
16 def __init__(self,out_dim=1,use_coadjoint=True):
17  super().__init__(); self.use_coadjoint=use_coadjoint; self.inp=nn.Linear(3,HIDDEN)
18  self.hamiltonian=nn.Sequential(nn.Linear(HIDDEN,96),nn.Tanh(),nn.Linear(96,1)); self.head=nn.Linear(HIDDEN,out_dim)
19 def vector_field(self,q,p,u):
20  with torch.enable_grad():
21   qq=q.requires_grad_(True); pp=p.requires_grad_(True)
22   H=self.hamiltonian(torch.cat([qq,pp],1)).sum()
23   gq,gp=torch.autograd.grad(H,(qq,pp),create_graph=self.training,retain_graph=self.training)
24   qdot=gp; pdot=-gq
25  if self.use_coadjoint:
26   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()
27   pd[:,:,0]+=p3[:,:,1]*v3[:,:,2]-p3[:,:,2]*v3[:,:,1]
28   pd[:,:,1]+=p3[:,:,2]*v3[:,:,0]-p3[:,:,0]*v3[:,:,2]
29   pd[:,:,2]+=p3[:,:,0]*v3[:,:,1]-p3[:,:,1]*v3[:,:,0]; pdot=pd.reshape(-1,HIDDEN//2)
30  return qdot,pdot
31 def transition(self,q,p,u):
32  qd,pd=self.vector_field(q,p,u); return q+qd,p+pd
33 def forward(self,x):
34  s=x.view(x.shape[0],-1,3); q,p=torch.chunk(torch.tanh(self.inp(s[:,0])),2,1)
35  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)
36  return self.head(torch.cat([q,p],1))
37def seed_all(s):
38 random.seed(s);np.random.seed(s);torch.manual_seed(s)
39 if torch.cuda.is_available():torch.cuda.manual_seed_all(s)
40def train(kind,cfg,seed,collect=False):
41 seed_all(seed); d=get_dataset('dynamics',seed,n_train=400,n_test=200); m=VanillaRNN() if kind=='base' else LiePoissonRNN()
42 net,metric,_=train_model(m,d,epochs=EPOCHS,lr=cfg['lr'],batch=BATCH,log=lambda *a:None)
43 if net is None:return (float('nan'),[],0) if collect else float('nan')
44 if not collect:return float(metric)
45 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=[]
46 with torch.enable_grad():
47  for a in [.25,.5,1.]:
48   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()))
49 return float(metric),vals,count_params(net)
50def ev(v):
51 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())}
52def main():
53 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=[]
54 for c in grid:runs.append({'cfg':c,'eval':ev([train('idea',c,s) for s in SEEDS])})
55 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])
56 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())}
57 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))
58if __name__=='__main__':main()