import sys, json, math from pathlib import Path 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, sweep_baseline, make_report SEEDS=tuple(range(8)) LRS=[1e-3,3e-3,1e-2] EPOCHS=15 BATCH=128 class ExplicitRNNPH(nn.Module): """Matched explicit recurrent baseline: same ports, width, and head.""" def __init__(self, input_dim, out_dim, hidden=64, dt=.25): super().__init__(); self.hidden=hidden; self.dt=dt self.inp=nn.Linear(input_dim, hidden) self.rawJ=nn.Parameter(torch.randn(hidden,hidden)*.03) self.B=nn.Linear(hidden, hidden, bias=False) self.head=nn.Linear(hidden,out_dim) def forward(self,x): seq=x.view(x.shape[0],-1,3); h=torch.zeros(x.shape[0],self.hidden,device=x.device,dtype=x.dtype) for u in seq.unbind(1): q=self.inp(u); h=h+self.dt*((h@self.rawJ.T)+self.B(q)) return self.head(h) class MidpointPH(nn.Module): """GRU-sized recurrent predictor with M=I and skew learned J. Input is an additive port forcing; midpoint is solved exactly for constant J. """ def __init__(self, input_dim, out_dim, hidden=64, dt=.25): super().__init__(); self.hidden=hidden; self.dt=dt self.inp=nn.Linear(input_dim, hidden) self.rawJ=nn.Parameter(torch.randn(hidden,hidden)*.03) self.B=nn.Linear(hidden, hidden, bias=False) self.head=nn.Linear(hidden,out_dim) self.register_buffer('I',torch.eye(hidden)) def matrix(self): return self.rawJ-self.rawJ.T def transition(self,h,u): J=self.matrix(); a=self.I-.5*self.dt*J; b=self.I+.5*self.dt*J # M=I, implicit midpoint: (I-dt J/2)h'=(I+dt J/2)h+dt B u return torch.linalg.solve(a, (b@h.T) + self.dt*self.B(u).T).T def forward(self,x): seq=x.view(x.shape[0],-1,3); h=torch.zeros(x.shape[0],self.hidden,device=x.device,dtype=x.dtype) for u in seq.unbind(1): h=self.transition(h,self.inp(u)) return self.head(h) def free_roll(self,h,steps=20): z=h J=self.matrix(); a=self.I-.5*self.dt*J; b=self.I+.5*self.dt*J vals=[] for _ in range(steps): z=torch.linalg.solve(a,b@z.T).T; vals.append(.5*(z*z).sum(1)) return torch.stack(vals,1) def sanity(): torch.manual_seed(123); d=7; A=torch.randn(d,d); J=A-A.T; dt=.37 I=torch.eye(d); z=torch.randn(5,d); h0=.5*(z*z).sum(1); a=I-dt*J/2; b=I+dt*J/2 dr=[] for _ in range(1000): z=torch.linalg.solve(a,b@z.T).T; dr.append(float((.5*(z*z).sum(1)-h0).abs().max())) return {'max_abs_energy_drift':max(dr),'skew_error':float(torch.linalg.norm(J+J.T)), 'predicted': 'zero up to floating point'} def make_base(cfg): def fn(seed): torch.manual_seed(seed); np.random.seed(seed); ds=get_dataset('dynamics',seed,n_train=4000,n_test=1000) net=ExplicitRNNPH(3,ds['ytr'].shape[1],64) _,metric,_=train_model(net,ds,epochs=EPOCHS,lr=cfg['lr'],batch=BATCH,log=lambda *_:None) return metric return fn def make_idea(cfg): def fn(seed): torch.manual_seed(seed); np.random.seed(seed); ds=get_dataset('dynamics',seed,n_train=4000,n_test=1000) net=MidpointPH(3,ds['ytr'].shape[1],64) _,metric,_=train_model(net,ds,epochs=EPOCHS,lr=cfg['lr'],batch=BATCH,log=lambda *_:None) return metric return fn def train_capture(cls, seed): torch.manual_seed(seed); ds=get_dataset('dynamics',seed,n_train=4000,n_test=1000) net=cls(ds); trained,metric,_=train_model(net,ds,epochs=EPOCHS,lr=3e-3,batch=BATCH,log=lambda *_:None) return trained,metric def main(): sanity_result=sanity() # The union of all idea learning rates is also the baseline grid. base=sweep_baseline(make_base,[{'lr':x} for x in LRS]) best_lr=base['best_cfg']['lr'] # Evaluate idea at best baseline lr and two nearby shared settings. idea_runs=[] for idea_lr in LRS: vals=[] for s in SEEDS: torch.manual_seed(s); np.random.seed(s); ds=get_dataset('dynamics',s,n_train=4000,n_test=1000) net=MidpointPH(3,ds['ytr'].shape[1],64) trained,m,_=train_model(net,ds,epochs=EPOCHS,lr=idea_lr,batch=BATCH,log=lambda *_:None) vals.append(float(m)) idea_runs.append({'lr':idea_lr,'mean':float(np.mean(vals)),'std':float(np.std(vals)),'per_seed':vals,'n':len(vals)}) chosen=min(idea_runs,key=lambda r:r['mean']) idea=dict(chosen); idea['sweep']=idea_runs; idea['cfg']={'lr':chosen['lr'],'nearby_tested':LRS} idea_lr=chosen['lr'] # Signature is measured on trained benchmark models, not the toy system. torch.manual_seed(0); ds=get_dataset('dynamics',0,n_train=4000,n_test=1000) im=MidpointPH(3,1,64); im,_,_=train_model(im,ds,epochs=EPOCHS,lr=idea_lr,batch=BATCH,log=lambda *_:None) with torch.no_grad(): h=torch.randn(32,64,device=next(im.parameters()).device); energies=im.free_roll(h,30); drift=float((energies-energies[:,0:1]).abs().max()) skew=float(torch.linalg.norm(im.matrix()+im.matrix().T)) report=make_report('dynamics','matched_recurrent_transition',base,idea,{'prediction':'unforced quadratic energy remains constant under midpoint skew transition','trained_model_observed_max_energy_drift':drift,'trained_model_skew_frobenius_error':skew,'predicted_energy_drift':0.0,'confirmed':bool(drift < 2e-4 and skew < 1e-6),'sanity':sanity_result,'note':'idea uses same input/output widths and hidden width; only recurrent transition differs'}) report['stage2_config']={'epochs':EPOCHS,'batch':BATCH,'lr_union':LRS,'seeds':list(SEEDS)} Path('bench_report.json').write_text(json.dumps(report,indent=2)) print(json.dumps(report,indent=2)) if __name__=='__main__': main()