Skew-Midpoint Neural Dynamics / bench_runner.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  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()