Skew-Midpoint Neural Dynamics / bench_runner.py
Mechanism confirmed, baseline not beaten
1import sys, json, math
2from pathlib import Path
3import numpy as np
4import torch
5import torch.nn as nn
6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
7from bench import get_dataset, make_model, train_model, sweep_baseline, make_report
8
9SEEDS=tuple(range(8))
10LRS=[1e-3,3e-3,1e-2]
11EPOCHS=15
12BATCH=128
13
14class ExplicitRNNPH(nn.Module):
15 """Matched explicit recurrent baseline: same ports, width, and head."""
16 def __init__(self, input_dim, out_dim, hidden=64, dt=.25):
17 super().__init__(); self.hidden=hidden; self.dt=dt
18 self.inp=nn.Linear(input_dim, hidden)
19 self.rawJ=nn.Parameter(torch.randn(hidden,hidden)*.03)
20 self.B=nn.Linear(hidden, hidden, bias=False)
21 self.head=nn.Linear(hidden,out_dim)
22 def forward(self,x):
23 seq=x.view(x.shape[0],-1,3); h=torch.zeros(x.shape[0],self.hidden,device=x.device,dtype=x.dtype)
24 for u in seq.unbind(1):
25 q=self.inp(u); h=h+self.dt*((h@self.rawJ.T)+self.B(q))
26 return self.head(h)
27
28class MidpointPH(nn.Module):
29 """GRU-sized recurrent predictor with M=I and skew learned J.
30 Input is an additive port forcing; midpoint is solved exactly for constant J.
31 """
32 def __init__(self, input_dim, out_dim, hidden=64, dt=.25):
33 super().__init__(); self.hidden=hidden; self.dt=dt
34 self.inp=nn.Linear(input_dim, hidden)
35 self.rawJ=nn.Parameter(torch.randn(hidden,hidden)*.03)
36 self.B=nn.Linear(hidden, hidden, bias=False)
37 self.head=nn.Linear(hidden,out_dim)
38 self.register_buffer('I',torch.eye(hidden))
39 def matrix(self): return self.rawJ-self.rawJ.T
40 def transition(self,h,u):
41 J=self.matrix(); a=self.I-.5*self.dt*J; b=self.I+.5*self.dt*J
42 # M=I, implicit midpoint: (I-dt J/2)h'=(I+dt J/2)h+dt B u
43 return torch.linalg.solve(a, (b@h.T) + self.dt*self.B(u).T).T
44 def forward(self,x):
45 seq=x.view(x.shape[0],-1,3); h=torch.zeros(x.shape[0],self.hidden,device=x.device,dtype=x.dtype)
46 for u in seq.unbind(1): h=self.transition(h,self.inp(u))
47 return self.head(h)
48 def free_roll(self,h,steps=20):
49 z=h
50 J=self.matrix(); a=self.I-.5*self.dt*J; b=self.I+.5*self.dt*J
51 vals=[]
52 for _ in range(steps):
53 z=torch.linalg.solve(a,b@z.T).T; vals.append(.5*(z*z).sum(1))
54 return torch.stack(vals,1)
55
56def sanity():
57 torch.manual_seed(123); d=7; A=torch.randn(d,d); J=A-A.T; dt=.37
58 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
59 dr=[]
60 for _ in range(1000):
61 z=torch.linalg.solve(a,b@z.T).T; dr.append(float((.5*(z*z).sum(1)-h0).abs().max()))
62 return {'max_abs_energy_drift':max(dr),'skew_error':float(torch.linalg.norm(J+J.T)), 'predicted': 'zero up to floating point'}
63
64def make_base(cfg):
65 def fn(seed):
66 torch.manual_seed(seed); np.random.seed(seed); ds=get_dataset('dynamics',seed,n_train=4000,n_test=1000)
67 net=ExplicitRNNPH(3,ds['ytr'].shape[1],64)
68 _,metric,_=train_model(net,ds,epochs=EPOCHS,lr=cfg['lr'],batch=BATCH,log=lambda *_:None)
69 return metric
70 return fn
71
72def make_idea(cfg):
73 def fn(seed):
74 torch.manual_seed(seed); np.random.seed(seed); ds=get_dataset('dynamics',seed,n_train=4000,n_test=1000)
75 net=MidpointPH(3,ds['ytr'].shape[1],64)
76 _,metric,_=train_model(net,ds,epochs=EPOCHS,lr=cfg['lr'],batch=BATCH,log=lambda *_:None)
77 return metric
78 return fn
79
80def train_capture(cls, seed):
81 torch.manual_seed(seed); ds=get_dataset('dynamics',seed,n_train=4000,n_test=1000)
82 net=cls(ds); trained,metric,_=train_model(net,ds,epochs=EPOCHS,lr=3e-3,batch=BATCH,log=lambda *_:None)
83 return trained,metric
84
85def main():
86 sanity_result=sanity()
87 # The union of all idea learning rates is also the baseline grid.
88 base=sweep_baseline(make_base,[{'lr':x} for x in LRS])
89 best_lr=base['best_cfg']['lr']
90 # Evaluate idea at best baseline lr and two nearby shared settings.
91 idea_runs=[]
92 for idea_lr in LRS:
93 vals=[]
94 for s in SEEDS:
95 torch.manual_seed(s); np.random.seed(s); ds=get_dataset('dynamics',s,n_train=4000,n_test=1000)
96 net=MidpointPH(3,ds['ytr'].shape[1],64)
97 trained,m,_=train_model(net,ds,epochs=EPOCHS,lr=idea_lr,batch=BATCH,log=lambda *_:None)
98 vals.append(float(m))
99 idea_runs.append({'lr':idea_lr,'mean':float(np.mean(vals)),'std':float(np.std(vals)),'per_seed':vals,'n':len(vals)})
100 chosen=min(idea_runs,key=lambda r:r['mean'])
101 idea=dict(chosen); idea['sweep']=idea_runs; idea['cfg']={'lr':chosen['lr'],'nearby_tested':LRS}
102 idea_lr=chosen['lr']
103 # Signature is measured on trained benchmark models, not the toy system.
104 torch.manual_seed(0); ds=get_dataset('dynamics',0,n_train=4000,n_test=1000)
105 im=MidpointPH(3,1,64); im,_,_=train_model(im,ds,epochs=EPOCHS,lr=idea_lr,batch=BATCH,log=lambda *_:None)
106 with torch.no_grad():
107 h=torch.randn(32,64,device=next(im.parameters()).device); energies=im.free_roll(h,30); drift=float((energies-energies[:,0:1]).abs().max())
108 skew=float(torch.linalg.norm(im.matrix()+im.matrix().T))
109 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'})
110 report['stage2_config']={'epochs':EPOCHS,'batch':BATCH,'lr_union':LRS,'seeds':list(SEEDS)}
111 Path('bench_report.json').write_text(json.dumps(report,indent=2))
112 print(json.dumps(report,indent=2))
113if __name__=='__main__': main()