import sys, json, random, numpy as np, torch import torch.nn as nn sys.path.insert(0, '/home/maxwelhelp/all/math2nn') from bench import get_dataset, train_model, evaluate, sweep_baseline, make_report SEEDS = tuple(range(8)) # Union used by both sides; baseline sweep includes all idea learning rates. GRID = [{'lr': 1e-3, 'epochs': 12}, {'lr': 3e-3, 'epochs': 12}, {'lr': 1e-2, 'epochs': 12}] class LiftRNN(nn.Module): """Finite m=2 skew-dilation recurrent cell with Cayley propagation.""" def __init__(self, input_dim=3, out_dim=1, hidden=64, eta=.05): super().__init__(); self.d=hidden; self.eta=eta self.rawS=nn.Parameter(torch.randn(hidden,hidden)*.02) self.rawA=nn.Parameter(torch.randn(hidden,hidden)*.02) self.inp=nn.Linear(input_dim, 2*hidden) self.head=nn.Linear(hidden,out_dim) def forward(self,x): x=x.view(x.shape[0],-1,3) B=x.shape[0]; dev=x.device S=self.rawS+self.rawS.T; A=self.rawA-self.rawA.T z0=torch.zeros((2*self.d,),device=dev,dtype=x.dtype) L=torch.zeros((2*self.d,2*self.d),device=dev,dtype=x.dtype) L[:self.d,:self.d]=A; L[:self.d,self.d:]=S L[self.d:,:self.d]=-S; L[self.d:,self.d:]=A I=torch.eye(2*self.d,device=dev,dtype=x.dtype) C=torch.linalg.solve(I-self.eta*L/2,I+self.eta*L/2) z=torch.zeros((B,2*self.d),device=dev,dtype=x.dtype) for k in range(x.shape[1]): z=z@C.T + self.inp(x[:,k]) # bounded nonlinear readout is only the task head; recurrent lift remains linear/unitary return self.head(z[:,:self.d]) def baseline_fn(cfg, seed): torch.manual_seed(seed); np.random.seed(seed); random.seed(seed) return make_base(3,1) def make_base(input_dim=3,out_dim=1): class Base(nn.Module): def __init__(self): super().__init__(); self.rnn=nn.GRU(input_dim,64,batch_first=True); self.head=nn.Linear(64,out_dim) def forward(self,x): _,h=self.rnn(x.view(x.shape[0],-1,input_dim)); return self.head(h[-1]) return Base() def idea_fn(cfg, seed): torch.manual_seed(seed); np.random.seed(seed); random.seed(seed) return LiftRNN(3,1,64,eta=.05) def run_one(fn, cfg, seed): ds=get_dataset('dynamics',seed,n_train=400,n_test=200) net,metric,hist=train_model(fn(cfg,seed),ds,epochs=cfg['epochs'],lr=cfg['lr'],batch=128) if net is None: raise RuntimeError('training failed') return metric, net, ds def eval_cfg(fn,cfg,seeds=SEEDS): vals=[] for s in seeds: vals.append(run_one(fn,cfg,s)[0]) return {'per_seed':vals,'mean':float(np.mean(vals)),'std':float(np.std(vals,ddof=1))} def signature(): # measured on a trained benchmark model: compare Cayley norm preservation to Euler growth metric,net,ds=run_one(idea_fn,GRID[1],0); net.eval(); d=net.d with torch.no_grad(): S=net.rawS+net.rawS.T; A=net.rawA-net.rawA.T L=torch.zeros((2*d,2*d)); L[:d,:d]=A; L[:d,d:]=S; L[d:,:d]=-S; L[d:,d:]=A I=torch.eye(2*d); eta=net.eta C=torch.linalg.solve(I-eta*L/2,I+eta*L/2); E=torch.randn(2*d) zc=E.clone(); ze=E.clone() for _ in range(20): zc=C@zc; ze=(I+eta*L)@ze cay=float((zc.norm()/E.norm()-1).abs()); euler=float(ze.norm()/E.norm()) return {'prediction':'Cayley preserves lifted norm; explicit Euler grows it', 'observed_cayley_relative_norm_error':cay, 'observed_euler_20_step_norm_ratio':euler, 'trained_test_mse':float(metric),'confirmed':bool(cay<1e-5 and euler>1.0001)} def main(): # baseline sweep is explicitly over the same union as the idea settings # Invoke the canonical harness sweep: make_fn(cfg) returns train_fn(seed). from bench import sweep_baseline def tuned_make(cfg): def train_fn(seed): return run_one(baseline_fn, cfg, int(seed))[0] return train_fn tuned=sweep_baseline(tuned_make, GRID, seeds=tuple(range(4))) sweep=[] for cfg in GRID: r=eval_cfg(baseline_fn,cfg); sweep.append({'cfg':cfg,**r}) best=min(sweep,key=lambda r:r['mean']) base={'best_cfg':best['cfg'],'sweep':sweep,'harness_tuning':tuned,'full':eval_cfg(baseline_fn,best['cfg'])} idea=[] for cfg in GRID: r=eval_cfg(idea_fn,cfg); idea.append({'cfg':cfg,**r}) ibest=min(idea,key=lambda r:r['mean']); idea_res={k:ibest[k] for k in ('per_seed','mean','std')} rep=make_report('dynamics','rnn_small',base,idea_res,{'idea_sweep':idea,'mechanism_signature':signature()}) rep['custom_track']=None with open('bench_report.json','w') as f: json.dump(rep,f,indent=2) print(json.dumps(rep,indent=2)) if __name__=='__main__': main()